diff --git a/src/openai/openai.service.spec.ts b/src/openai/openai.service.spec.ts new file mode 100644 index 0000000..d3b35a1 --- /dev/null +++ b/src/openai/openai.service.spec.ts @@ -0,0 +1,141 @@ +import { firstValueFrom, toArray } from 'rxjs'; +import { AzureOpenAI } from 'openai'; +import { OpenaiService } from './openai.service'; +import { AppConfigService } from '../common/config/appConfig.service'; + +jest.mock('openai', () => ({ + __esModule: true, + default: jest.fn(), + AzureOpenAI: jest.fn(), +})); + +const CONFIG: Record = { + aiProvider: 'openai-azure', + openaiAzureEndpoint: 'https://example.openai.azure.com', + openaiAzureKey: 'key', + openaiAzureVersion: '2024-06-01', +}; + +const contentChunk = (content: string) => ({ + choices: [{ index: 0, finish_reason: null, delta: { content } }], +}); + +async function* streamOf(chunks: unknown[], error?: Error) { + for (const chunk of chunks) yield chunk; + if (error) throw error; +} + +describe('OpenaiService', () => { + let service: OpenaiService; + let create: jest.Mock; + + beforeEach(() => { + create = jest.fn(); + (AzureOpenAI as unknown as jest.Mock).mockImplementation(() => ({ + apiKey: 'key', + chat: { completions: { create } }, + })); + + service = new OpenaiService({ + get: (key: string) => CONFIG[key], + } as unknown as AppConfigService); + + jest.spyOn(service, 'analyzeChatConversation').mockResolvedValue(''); + jest.spyOn(service, 'implementApiCalls').mockResolvedValue(undefined); + jest.spyOn(service['logger'], 'error').mockImplementation(() => undefined); + }); + + const requestStream = (completeCb?: jest.Mock) => + service.getChatGptCompletionStream( + { + messages: [{ role: 'user', content: 'Hi' }], + model: 'gpt-4o', + stream: true, + }, + completeCb, + ); + + describe('getChatGptCompletion', () => { + it.each([ + ['no choices', []], + [ + 'a filtered choice without message', + [{ index: 0, finish_reason: 'content_filter' }], + ], + [ + 'a null message content', + [{ index: 0, message: { role: 'assistant', content: null } }], + ], + ])('returns an empty response for %s', async (_, choices) => { + create.mockResolvedValue({ choices, usage: undefined }); + + const result = await service.getChatGptCompletion({ + messages: [{ role: 'user', content: 'Hi' }], + model: 'gpt-4o', + }); + + expect(result.response).toBe(''); + }); + }); + + describe('getChatGptCompletionStream', () => { + it('skips chunks without delta, such as Azure content filter chunks', async () => { + create.mockResolvedValue( + streamOf([ + { choices: [] }, + { + choices: [ + { index: 0, finish_reason: null, content_filter_results: {} }, + ], + }, + contentChunk('Hello'), + { choices: [{ index: 0, finish_reason: null, delta: {} }] }, + contentChunk(' world'), + { choices: [{ index: 0, finish_reason: 'stop', delta: {} }] }, + ]), + ); + const completeCb = jest.fn().mockResolvedValue(undefined); + + const observable = await requestStream(completeCb); + const values = await firstValueFrom(observable.pipe(toArray())); + + expect(values).toEqual([ + JSON.stringify({ content: 'Hello' }), + JSON.stringify({ content: ' world' }), + '[DONE]', + ]); + expect(completeCb).toHaveBeenCalledWith( + 'Hello world', + expect.objectContaining({ prompt: expect.any(Number) }), + ); + }); + + it('errors the observable instead of rejecting when the stream fails', async () => { + create.mockResolvedValue( + streamOf([contentChunk('Hel')], new Error('connection reset')), + ); + const completeCb = jest.fn(); + + const observable = await requestStream(completeCb); + + await expect(firstValueFrom(observable.pipe(toArray()))).rejects.toThrow( + 'Failed to generate answer', + ); + expect(completeCb).not.toHaveBeenCalled(); + }); + + it('logs instead of rejecting when the completion callback fails', async () => { + create.mockResolvedValue(streamOf([contentChunk('Hi')])); + const completeCb = jest.fn().mockRejectedValue(new Error('db down')); + + const observable = await requestStream(completeCb); + await firstValueFrom(observable.pipe(toArray())); + await new Promise(process.nextTick); + + expect(service['logger'].error).toHaveBeenCalledWith( + expect.stringContaining('callback'), + expect.any(Error), + ); + }); + }); +}); diff --git a/src/openai/openai.service.ts b/src/openai/openai.service.ts index 47ce7e8..ec052d8 100644 --- a/src/openai/openai.service.ts +++ b/src/openai/openai.service.ts @@ -344,7 +344,7 @@ export class OpenaiService { // API Call try { const res = await openAiClient.chat.completions.create(data); - const chatResponse = res.choices[0].message.content; + const chatResponse = res.choices[0]?.message?.content ?? ''; return { response: chatResponse, @@ -402,43 +402,70 @@ export class OpenaiService { data.messages.map((m) => m.content).join(' '), ); + let completionStream: AsyncIterable; try { - const completionStream = await openAiClient.chat.completions.create(data); + completionStream = await openAiClient.chat.completions.create(data); + } catch (error) { + if (error instanceof APIError) { + this.logger.error('OpenAI ChatCompletion API error', error); + this.logger.error('Error response', error.error); + } + throw error; + } - let answer = ''; + // Not awaited: the caller needs the observable before chunks arrive. + // forwardCompletionStream never rejects. + void this.forwardCompletionStream( + completionStream, + observable, + promptTokens, + completeCb, + ); - const streamPromise = new Promise(async (res) => { - for await (const part of completionStream) { - if (part.choices.length === 0) continue; + return observable; + } + + private async forwardCompletionStream( + completionStream: AsyncIterable, + observable: Subject, + promptTokens: number, + completeCb?: ( + answer: string, + usage: ChatGTPResponse['tokenUsage'], + ) => Promise, + ) { + let answer = ''; - const { content } = part.choices[0].delta; + try { + for await (const part of completionStream) { + // Azure sends chunks without `delta` (e.g. content filter results) + const content = part.choices?.[0]?.delta?.content; + if (content == null) continue; - if (content !== undefined) { - observable.next(JSON.stringify({ content })); - answer += content; - } - } + observable.next(JSON.stringify({ content })); + answer += content; + } + } catch (error) { + this.logger.error('OpenAI ChatCompletion stream error', error); + observable.error(new Error('Failed to generate answer')); + return; + } - res(true); - }); + observable.next('[DONE]'); + observable.complete(); - streamPromise.then(() => { - observable.next('[DONE]'); - observable.complete(); - const completionTokens = this.getTokenCount(answer); - completeCb?.(answer, { - prompt: promptTokens, - completion: completionTokens, - total: promptTokens + completionTokens, - }); + try { + const completionTokens = this.getTokenCount(answer); + await completeCb?.(answer, { + prompt: promptTokens, + completion: completionTokens, + total: promptTokens + completionTokens, }); } catch (error) { - if (APIError.isPrototypeOf(error)) { - this.logger.error('OpenAI ChatCompletion API error', error); - this.logger.error('Error response', error.data); - } - throw error; + this.logger.error( + 'OpenAI ChatCompletion completion callback error', + error, + ); } - return observable; } }