155 lines
5.6 KiB
TypeScript
155 lines
5.6 KiB
TypeScript
import { ConflictException, NotFoundException } from '@nestjs/common';
|
|
import { AiChatService } from './ai-chat.service';
|
|
|
|
const authenticatedUser = {
|
|
id: 7,
|
|
username: 'tester',
|
|
permissions: ['ai:chat:use'],
|
|
isSuperAdmin: false,
|
|
};
|
|
|
|
function createService(conversationOverrides: Record<string, unknown> = {}) {
|
|
const conversations = {
|
|
findOne: jest.fn(),
|
|
find: jest.fn(),
|
|
create: jest.fn((value) => value),
|
|
save: jest.fn(async (value) => ({ id: 1, ...value })),
|
|
remove: jest.fn(),
|
|
...conversationOverrides,
|
|
};
|
|
const service = new AiChatService(
|
|
conversations as never,
|
|
{ exists: jest.fn().mockResolvedValue(false) } as never,
|
|
{} as never,
|
|
{} as never,
|
|
{} as never,
|
|
{} as never,
|
|
{} as never,
|
|
);
|
|
return { service, conversations };
|
|
}
|
|
|
|
describe('AiChatService', () => {
|
|
it('按 userId 查询会话,无法借 id 访问其他用户会话', async () => {
|
|
const { service, conversations } = createService({ findOne: jest.fn().mockResolvedValue(null) });
|
|
await expect(service.getMessages(7, 99)).rejects.toBeInstanceOf(NotFoundException);
|
|
expect(conversations.findOne).toHaveBeenCalledWith({ where: { id: 99, userId: 7 } });
|
|
});
|
|
|
|
it('生成中的会话禁止删除', async () => {
|
|
const entity = { id: 2, userId: 7 };
|
|
const { service, conversations } = createService({ findOne: jest.fn().mockResolvedValue(entity) });
|
|
(service as unknown as { activeConversations: Set<number> }).activeConversations.add(2);
|
|
await expect(service.deleteConversation(7, 2)).rejects.toBeInstanceOf(ConflictException);
|
|
expect(conversations.remove).not.toHaveBeenCalled();
|
|
});
|
|
|
|
it('并发获取同一会话时只允许一个请求进入生成流程', async () => {
|
|
let resolveExists!: (value: boolean) => void;
|
|
const exists = jest.fn(
|
|
() => new Promise<boolean>((resolve) => {
|
|
resolveExists = resolve;
|
|
}),
|
|
);
|
|
const { service } = createService();
|
|
(service as unknown as { messages: { exists: typeof exists } }).messages.exists = exists;
|
|
const acquire = (service as unknown as { acquireConversation(id: number): Promise<void> })
|
|
.acquireConversation.bind(service);
|
|
|
|
const first = acquire(5);
|
|
await expect(acquire(5)).rejects.toBeInstanceOf(ConflictException);
|
|
resolveExists(false);
|
|
await expect(first).resolves.toBeUndefined();
|
|
});
|
|
|
|
it('工具摘要脱敏并限制长度', () => {
|
|
const { service } = createService();
|
|
const summarize = (service as unknown as { summarize(value: unknown): string }).summarize.bind(service);
|
|
const summary = summarize({
|
|
phone: '13800138000',
|
|
idCard: '11010519491231002X',
|
|
note: `联系电话 13900139000 ${'x'.repeat(3000)}`,
|
|
apiKey: 'sk-sensitive-value',
|
|
});
|
|
expect(summary).not.toContain('13800138000');
|
|
expect(summary).not.toContain('13900139000');
|
|
expect(summary).not.toContain('11010519491231002X');
|
|
expect(summary).not.toContain('sk-sensitive-value');
|
|
expect(summary.length).toBeLessThanOrEqual(2000);
|
|
});
|
|
|
|
it.each([
|
|
{ abort: false, expectedStatus: 'failed', expectedCode: 'UPSTREAM_ERROR' },
|
|
{ abort: true, expectedStatus: 'cancelled', expectedCode: 'CLIENT_ABORTED' },
|
|
])('流中断后保存已生成内容和 $expectedStatus 状态', async ({ abort, expectedStatus, expectedCode }) => {
|
|
const conversation = { id: 3, userId: 7, title: '测试', lastMessageAt: null };
|
|
const assistant = {
|
|
id: 12,
|
|
conversationId: 3,
|
|
role: 'assistant',
|
|
content: '',
|
|
reasoningContent: null,
|
|
status: 'pending',
|
|
errorCode: null,
|
|
};
|
|
const messageSave = jest.fn(async (value) => value);
|
|
const messages = {
|
|
exists: jest.fn().mockResolvedValue(false),
|
|
find: jest.fn().mockResolvedValue([]),
|
|
save: messageSave,
|
|
};
|
|
const manager = {
|
|
create: jest.fn((_entity, value) => value),
|
|
save: jest
|
|
.fn()
|
|
.mockResolvedValueOnce({ id: 11, conversationId: 3, role: 'user', content: '查询' })
|
|
.mockResolvedValueOnce(assistant),
|
|
update: jest.fn(),
|
|
};
|
|
const abortController = new AbortController();
|
|
const modelStream = {
|
|
stream: async function* () {
|
|
yield { type: 'content' as const, delta: '部分回答' };
|
|
if (abort) {
|
|
abortController.abort(new Error('client disconnected'));
|
|
yield { type: 'complete' as const, toolCalls: [] };
|
|
return;
|
|
}
|
|
throw new Error('upstream failed');
|
|
},
|
|
};
|
|
const service = new AiChatService(
|
|
{ findOne: jest.fn().mockResolvedValue(conversation) } as never,
|
|
messages as never,
|
|
{ save: jest.fn() } as never,
|
|
{ transaction: jest.fn(async (callback) => callback(manager)) } as never,
|
|
{ getRuntimeConfig: jest.fn().mockResolvedValue({}) } as never,
|
|
{ listAvailable: jest.fn().mockReturnValue([]) } as never,
|
|
modelStream as never,
|
|
);
|
|
const emitted: Array<{ event: string; data: Record<string, unknown> }> = [];
|
|
const run = service.streamMessage(
|
|
authenticatedUser as never,
|
|
3,
|
|
'查询',
|
|
abortController.signal,
|
|
(event, data) => emitted.push({ event, data }),
|
|
jest.fn(),
|
|
);
|
|
|
|
if (abort) await expect(run).resolves.toBeUndefined();
|
|
else await expect(run).rejects.toThrow('upstream failed');
|
|
|
|
expect(messageSave).toHaveBeenCalledWith(
|
|
expect.objectContaining({
|
|
id: 12,
|
|
content: '部分回答',
|
|
status: expectedStatus,
|
|
errorCode: expectedCode,
|
|
}),
|
|
);
|
|
expect(emitted.some(({ event }) => event === 'content.delta')).toBe(true);
|
|
expect(emitted.some(({ event }) => event === 'message.cancelled')).toBe(abort);
|
|
});
|
|
});
|