Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
141 changes: 141 additions & 0 deletions src/openai/openai.service.spec.ts
Original file line number Diff line number Diff line change
@@ -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<string, string> = {
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),
);
});
});
});
85 changes: 56 additions & 29 deletions src/openai/openai.service.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -402,43 +402,70 @@ export class OpenaiService {
data.messages.map((m) => m.content).join(' '),
);

let completionStream: AsyncIterable<OpenAI.Chat.ChatCompletionChunk>;
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<OpenAI.Chat.ChatCompletionChunk>,
observable: Subject<string>,
promptTokens: number,
completeCb?: (
answer: string,
usage: ChatGTPResponse['tokenUsage'],
) => Promise<void>,
) {
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;
}
}
Loading