From bd3667c6e86450689aacffaa3e2c8049e859ebbb Mon Sep 17 00:00:00 2001 From: Reversean Date: Wed, 29 Jul 2026 19:34:51 +0300 Subject: [PATCH] feat(ai): stream AI suggestions over HTTP An Ask AI suggestion takes tens of seconds to generate, and the GraphQL resolver can only return it once the model has finished. Suggestions are now also available as a stream over a plain Express route, guarded by the same workspace membership check as the resolver and rejecting requests missing the project, event or repetition id before reaching the model. --- .eslintrc.js | 8 + package.json | 2 +- src/directives/requireUserInWorkspace.ts | 2 +- src/index.ts | 6 + src/integrations/vercel-ai/index.ts | 37 ++++- src/integrations/vercel-ai/routes.ts | 120 ++++++++++++++ src/services/ai.ts | 61 ++++++- src/services/types.ts | 4 +- test/helpers/expressRequest.ts | 144 ++++++++++++++++ test/integrations/ai-routes.test.ts | 200 +++++++++++++++++++++++ test/integrations/github-routes.test.ts | 124 +++----------- test/integrations/vercel-ai.test.ts | 24 ++- test/services/askAi.test.ts | 27 ++- 13 files changed, 639 insertions(+), 120 deletions(-) create mode 100644 src/integrations/vercel-ai/routes.ts create mode 100644 test/helpers/expressRequest.ts create mode 100644 test/integrations/ai-routes.test.ts diff --git a/.eslintrc.js b/.eslintrc.js index 12245e12c..7cec6e3fc 100644 --- a/.eslintrc.js +++ b/.eslintrc.js @@ -4,6 +4,14 @@ module.exports = { 'node': true, 'jest': true }, + globals: { + /** + * WHATWG Fetch/Streams API globals available in Node 18+ (this project runs on Node 24 + * per .nvmrc) - not part of eslint's "node" env, which predates them + */ + 'ReadableStream': 'readonly', + 'Response': 'readonly' + }, rules: { '@typescript-eslint/camelcase': 'warn', '@typescript-eslint/no-unused-vars': 'warn', diff --git a/package.json b/package.json index a310be750..638a7e30a 100644 --- a/package.json +++ b/package.json @@ -1,6 +1,6 @@ { "name": "hawk.api", - "version": "1.5.11", + "version": "1.5.12", "main": "index.ts", "license": "BUSL-1.1", "scripts": { diff --git a/src/directives/requireUserInWorkspace.ts b/src/directives/requireUserInWorkspace.ts index 092b651b0..1626cccd4 100644 --- a/src/directives/requireUserInWorkspace.ts +++ b/src/directives/requireUserInWorkspace.ts @@ -37,7 +37,7 @@ async function checkUserInWorkspaceByWorkspaceId(context: ResolverContextBase, w * @param context - request context * @param projectId - project id */ -async function checkUserInWorkspaceByProjectId(context: ResolverContextBase, projectId: string): Promise { +export async function checkUserInWorkspaceByProjectId(context: ResolverContextBase, projectId: string): Promise { const userId = context.user.id; if (userId) { diff --git a/src/index.ts b/src/index.ts index cb6f8d934..897d2c94b 100644 --- a/src/index.ts +++ b/src/index.ts @@ -32,6 +32,7 @@ import ReleasesFactory from './models/releasesFactory'; import RedisHelper from './redisHelper'; import { appendSsoRoutes } from './sso'; import { appendGitHubRoutes } from './integrations/github'; +import { appendAiAssistantRoutes } from './integrations/vercel-ai/routes'; /** * Option to enable playground @@ -272,6 +273,11 @@ class HawkAPI { */ appendGitHubRoutes(this.app, sharedFactories); + /** + * Append AI assistant route to Express app + */ + appendAiAssistantRoutes(this.app); + await this.server.start(); this.app.use(graphqlUploadExpress()); this.server.applyMiddleware({ app: this.app }); diff --git a/src/integrations/vercel-ai/index.ts b/src/integrations/vercel-ai/index.ts index 745e2624a..88812e7e3 100644 --- a/src/integrations/vercel-ai/index.ts +++ b/src/integrations/vercel-ai/index.ts @@ -1,4 +1,5 @@ -import { generateText } from 'ai'; +import { generateText, streamText } from 'ai'; +import { ProviderOptions } from '@ai-sdk/provider-utils'; /** * Params for a single completion call to the model @@ -29,11 +30,24 @@ class VercelAIApi { */ private readonly modelId: string; + /** + * Provider Gateway fallback order + */ + private readonly providerOptions: ProviderOptions; + + /** + * Set up model id and provider fallback order + */ constructor() { /** * @todo make it dynamic, get from project settings */ this.modelId = 'deepseek/deepseek-v4-flash'; + this.providerOptions = { + gateway: { + order: ['novita', 'azure', 'deepseek'], + }, + }; } /** @@ -47,15 +61,26 @@ class VercelAIApi { model: this.modelId, system, prompt, - providerOptions: { - gateway: { - order: ['novita', 'azure', 'deepseek'], - }, - }, + providerOptions: this.providerOptions, }); return text; } + + /** + * Send a system/prompt pair to the model and return the generated text as a stream + * + * @param {CompletionParams} params - system instruction and prompt to complete + * @returns {StreamTextResult} text generated by the model, as a stream + */ + public stream({ system, prompt }: CompletionParams): ReturnType { + return streamText({ + model: this.modelId, + system, + prompt, + providerOptions: this.providerOptions, + }); + } } export const vercelAIApi = new VercelAIApi(); diff --git a/src/integrations/vercel-ai/routes.ts b/src/integrations/vercel-ai/routes.ts new file mode 100644 index 000000000..454f4f84a --- /dev/null +++ b/src/integrations/vercel-ai/routes.ts @@ -0,0 +1,120 @@ +import '../../typeDefs/expressContext'; +import express from 'express'; +import { Readable } from 'stream'; +import type { ReadableStream as NodeReadableStream } from 'stream/web'; +import { getEventsFactory } from '../../resolvers/helpers/eventsFactory'; +import { checkUserInWorkspaceByProjectId } from '../../directives/requireUserInWorkspace'; +import { aiService } from '../../services/ai'; + +/** + * Verify the requesting user is a member of the project's workspace. + * + * @param req - Express request + * @param res - Express response + * @param projectId - project id from query parameters + * @returns user ID if authorized, {@code null} otherwise (response already sent) + */ +async function authorizeProjectAccess( + req: express.Request, + res: express.Response, + projectId: string | undefined +): Promise { + const userId = req.context?.user?.id; + + if (!userId) { + res.status(401).json({ error: 'Unauthorized. Please provide authorization token.' }); + + return null; + } + + if (!projectId) { + res.status(400).json({ error: 'projectId query parameter is required' }); + + return null; + } + + try { + await checkUserInWorkspaceByProjectId(req.context, projectId); + } catch (error) { + res.status(403).json({ error: error instanceof Error ? error.message : 'You have no access to this workspace' }); + + return null; + } + + return userId; +} + +/** + * Create AI assistant router + * + * @returns Express router with AI assistant endpoints + */ +export function createAiStreamRouter(): express.Router { + const router = express.Router(); + + /** + * GET /integration/ai/stream?projectId=&eventId=&originalEventId= + * Stream an AI suggestion for the event + */ + router.get('/stream', async (req, res, next) => { + try { + const { projectId, eventId, originalEventId } = req.query; + + const userId = await authorizeProjectAccess(req, res, projectId as string | undefined); + + if (!userId) { + return; + } + + if (!eventId || typeof eventId !== 'string') { + res.status(400).json({ error: 'eventId query parameter is required' }); + + return; + } + + if (!originalEventId || typeof originalEventId !== 'string') { + res.status(400).json({ error: 'originalEventId query parameter is required' }); + + return; + } + + const eventsFactory = getEventsFactory(req.context, projectId as string); + + let result; + + try { + result = await aiService.streamSuggestion(eventsFactory, eventId, originalEventId); + } catch (error) { + res.status(404).json({ error: error instanceof Error ? error.message : 'Event not found' }); + + return; + } + + const response = result.toUIMessageStreamResponse(); + + res.status(response.status); + response.headers.forEach((value, key) => res.setHeader(key, value)); + + if (!response.body) { + res.end(); + + return; + } + + Readable.fromWeb(response.body as NodeReadableStream).pipe(res); + } catch (error) { + next(error); + } + }); + + return router; +} + +/** + * Append AI assistant routes to Express app + * + * @param app - Express application instance + */ +export function appendAiAssistantRoutes(app: express.Application): void { + app.use('/integration/ai', createAiStreamRouter()); +} diff --git a/src/services/ai.ts b/src/services/ai.ts index 5924cb5ff..ed9d689bd 100644 --- a/src/services/ai.ts +++ b/src/services/ai.ts @@ -4,6 +4,7 @@ import { buildEventPrompt, spotlightInstruction } from './askAi/security/spotlig import { isLeaked, SUGGESTION_FALLBACK_MESSAGE } from './askAi/security/leakDetector'; import { ctoInstruction } from './askAi/instructions/cto'; import { EventsFactoryInterface } from './types'; +import type { Event } from './types'; /** * Report that the leak tripwire fired. @@ -43,12 +44,12 @@ export class AIService { * @param originalEventId - original event id * @returns {Promise} - suggestion */ - public async generateSuggestion(eventsFactory: EventsFactoryInterface, eventId: string, originalEventId: string): Promise { - const event = await eventsFactory.getEventRepetition(eventId, originalEventId); - - if (!event) { - throw new Error('Event not found'); - } + public async generateSuggestion( + eventsFactory: EventsFactoryInterface, + eventId: string, + originalEventId: string + ): Promise { + const event = await this.getEventOrThrow(eventsFactory, eventId, originalEventId); const { prompt, nonce } = buildEventPrompt(event.payload); @@ -65,6 +66,54 @@ export class AIService { return text; } + + /** + * Generate streaming suggestion for the event + * + * The payload is spotlighted by {@link buildEventPrompt} exactly as in + * {@link AIService.generateSuggestion}. + * + * @param eventsFactory - events factory + * @param eventId - event id + * @param originalEventId - original event id + * @returns streaming suggestion + */ + public async streamSuggestion( + eventsFactory: EventsFactoryInterface, + eventId: string, + originalEventId: string + ): Promise> { + const event = await this.getEventOrThrow(eventsFactory, eventId, originalEventId); + + const { prompt, nonce } = buildEventPrompt(event.payload); + + return vercelAIApi.stream({ + system: ctoInstruction + spotlightInstruction(nonce), + prompt, + }); + } + + /** + * Find the event repetition or throw if it doesn't exist + * + * @param eventsFactory - events factory + * @param eventId - event id + * @param originalEventId - original event id + * @returns {Promise} - event repetition + */ + private async getEventOrThrow( + eventsFactory: EventsFactoryInterface, + eventId: string, + originalEventId: string + ): Promise { + const event = await eventsFactory.getEventRepetition(eventId, originalEventId); + + if (!event) { + throw new Error('Event not found'); + } + + return event; + } } export const aiService = new AIService(); diff --git a/src/services/types.ts b/src/services/types.ts index 1b14501fa..2007767b5 100644 --- a/src/services/types.ts +++ b/src/services/types.ts @@ -3,7 +3,7 @@ import { EventAddons, EventData } from '@hawk.so/types'; /** * Event type which is returned by events factory */ -type Event = { +export type Event = { _id: string; payload: EventData; }; @@ -20,4 +20,4 @@ export interface EventsFactoryInterface { * @returns {Promise>} - event repetition */ getEventRepetition(repetitionId: string, originalEventId: string): Promise; -} \ No newline at end of file +} diff --git a/test/helpers/expressRequest.ts b/test/helpers/expressRequest.ts new file mode 100644 index 000000000..7e31bbfe7 --- /dev/null +++ b/test/helpers/expressRequest.ts @@ -0,0 +1,144 @@ +import { Writable } from 'stream'; +import express from 'express'; + +export interface CapturedResponse { + status: number; + headers: Record; + body: any; +} + +/** + * Express's expressInit middleware unconditionally runs setPrototypeOf(res, app.response) + * on every request. That silently discards any *class* methods on our fake res (they live + * on the class prototype, not as own properties of the instance) and falls back to Express/ + * Node's real ServerResponse implementation, which then throws trying to touch a real socket + * that doesn't exist here. Own properties always shadow whatever a new prototype provides, + * so binding inherited methods as own properties makes them survive the prototype swap. + * + * @param obj - object whose inherited methods should survive a prototype swap + */ +function pinInheritedMethodsAsOwnProperties(obj: any): void { + let proto = Object.getPrototypeOf(obj); + + while (proto && proto !== Object.prototype) { + for (const key of Object.getOwnPropertyNames(proto)) { + if (key === 'constructor' || Object.prototype.hasOwnProperty.call(obj, key)) { + continue; + } + + const descriptor = Object.getOwnPropertyDescriptor(proto, key); + + if (descriptor && typeof descriptor.value === 'function') { + obj[key] = descriptor.value.bind(obj); + } + } + + proto = Object.getPrototypeOf(proto); + } +} + +/** + * Fake Express response supporting both res.json()/res.send()/res.redirect() and + * res.write()/res.end() via stream.pipe() (as the AI stream route does). Built on a real + * Writable so pipe() gets genuine EventEmitter semantics, with every method pinned as an + * own property per pinInheritedMethodsAsOwnProperties above. + * + * @param settle - called once with everything the route wrote to the response + * @returns {any} fake response object to hand to Express + */ +function createFakeResponse(settle: (result: CapturedResponse) => void): any { + let statusCode = 200; + const headers: Record = {}; + const chunks: Buffer[] = []; + let settled = false; + + function finish(body: any): void { + if (settled) { + return; + } + + settled = true; + settle({ + status: statusCode, + headers, + body, + }); + } + + const res: any = new Writable({ + write(chunk, _encoding, callback) { + chunks.push(Buffer.isBuffer(chunk) ? chunk : Buffer.from(chunk)); + callback(); + }, + final(callback) { + finish(Buffer.concat(chunks).toString('utf-8')); + callback(); + }, + }); + + pinInheritedMethodsAsOwnProperties(res); + + res.status = (code: number): any => { + statusCode = code; + + return res; + }; + res.setHeader = (key: string, value: string): any => { + headers[key] = value; + + return res; + }; + res.getHeader = (key: string): string | undefined => headers[key]; + res.json = (data: any): void => finish(data); + res.send = (data?: any): void => { + if (!settled) { + finish(data); + } + }; + res.redirect = (url: string): void => { + statusCode = 302; + finish(url); + }; + + return res; +} + +/** + * Sends a fake request through an Express app via its internal handle() method, + * without opening a real socket - simulates how Express actually processes requests. + * + * @param app - Express application to route the request through + * @param method - HTTP method + * @param path - request path, without the query string + * @param query - query parameters to append + * @returns {Promise} status, headers and body the route produced + */ +export function makeExpressRequest( + app: express.Application, + method: string, + path: string, + query?: Record +): Promise { + return new Promise((resolve, reject) => { + const url = query ? `${path}?${new URLSearchParams(query).toString()}` : path; + const req = { + method, + url, + originalUrl: url, + path, + query: query || {}, + headers: {}, + get: jest.fn(), + params: {}, + body: {}, + } as any; + + const res = createFakeResponse(resolve); + + (app as any).handle(req, res, (err: any) => { + if (err) { + reject(err); + } + }); + }); +} diff --git a/test/integrations/ai-routes.test.ts b/test/integrations/ai-routes.test.ts new file mode 100644 index 000000000..5f584bcea --- /dev/null +++ b/test/integrations/ai-routes.test.ts @@ -0,0 +1,200 @@ +import '../../src/env-test'; +import express from 'express'; +import { makeExpressRequest } from '../helpers/expressRequest'; + +import { aiService } from '../../src/services/ai'; +import { getEventsFactory } from '../../src/resolvers/helpers/eventsFactory'; +import { checkUserInWorkspaceByProjectId } from '../../src/directives/requireUserInWorkspace'; +import { createAiStreamRouter } from '../../src/integrations/vercel-ai/routes'; + +jest.mock('../../src/services/ai', () => ({ + aiService: { + streamSuggestion: jest.fn(), + }, +})); + +jest.mock('../../src/resolvers/helpers/eventsFactory', () => ({ + getEventsFactory: jest.fn(), +})); + +jest.mock('../../src/directives/requireUserInWorkspace', () => ({ + checkUserInWorkspaceByProjectId: jest.fn(), +})); + +const mockStreamSuggestion = aiService.streamSuggestion as jest.Mock; +const mockGetEventsFactory = getEventsFactory as jest.Mock; +const mockCheckUserInWorkspaceByProjectId = checkUserInWorkspaceByProjectId as jest.Mock; + +const userId = '507f1f77bcf86cd799439011'; +const projectId = '507f1f77bcf86cd799439022'; +const eventId = 'event-1'; +const originalEventId = 'original-event-1'; + +function setupApp(contextOverrides?: (req: any) => void): express.Application { + const app = express(); + + app.use((req: any, _res, next) => { + req.context = { + user: { id: userId }, + factories: {} as any, + }; + + if (contextOverrides) { + contextOverrides(req); + } + + next(); + }); + + app.use('/integration/ai', createAiStreamRouter()); + + return app; +} + +/** + * Builds a fake SSE Response matching what streamSuggestion(...).toUIMessageStreamResponse() + * really returns. start/start-step/reasoning-* chunks and the [DONE] terminator are copied + * verbatim from a live Vercel AI Gateway call; the remaining tail (text-* and finish-* chunks) + * follows the same envelope, per the ai@5.0.89 UI Message Stream Protocol types. + */ +function createFakeStreamResponse(): Response { + const chunks = [ + '{"type":"start"}', + '{"type":"start-step"}', + '{"type":"reasoning-start","id":"reasoning-0"}', + '{"type":"reasoning-delta","id":"reasoning-0","delta":"Reasoning"}', + '{"type":"reasoning-end","id":"reasoning-0"}', + '{"type":"text-start","id":"text-0"}', + '{"type":"text-delta","id":"text-0","delta":"Answer"}', + '{"type":"text-end","id":"text-0"}', + '{"type":"finish-step"}', + '{"type":"finish"}', + ]; + + const encoder = new TextEncoder(); + const body = new ReadableStream({ + start(controller) { + for (const chunk of chunks) { + controller.enqueue(encoder.encode(`data: ${chunk}\n\n`)); + } + controller.enqueue(encoder.encode('data: [DONE]\n\n')); + controller.close(); + }, + }); + + return new Response(body, { + status: 200, + headers: { + 'content-type': 'text/event-stream', + 'cache-control': 'no-cache', + connection: 'keep-alive', + 'x-vercel-ai-ui-message-stream': 'v1', + }, + }); +} + +describe('AI stream routes - GET /integration/ai/stream', () => { + beforeEach(() => { + jest.clearAllMocks(); + mockGetEventsFactory.mockReturnValue({}); + mockCheckUserInWorkspaceByProjectId.mockResolvedValue(undefined); + }); + + it('should return 401 when the user is not authenticated', async () => { + const app = setupApp((req) => { + req.context.user.id = undefined; + }); + + const response = await makeExpressRequest(app, 'GET', '/integration/ai/stream', { + projectId, + eventId, + originalEventId, + }); + + expect(response.status).toBe(401); + expect(response.body.error).toContain('Unauthorized'); + }); + + it('should return 400 when projectId is missing', async () => { + const app = setupApp(); + + const response = await makeExpressRequest(app, 'GET', '/integration/ai/stream', { + eventId, + originalEventId, + }); + + expect(response.status).toBe(400); + expect(response.body.error).toContain('projectId'); + }); + + it('should return 403 when the user has no access to the project workspace', async () => { + mockCheckUserInWorkspaceByProjectId.mockRejectedValue(new Error('You have no access to this workspace')); + const app = setupApp(); + + const response = await makeExpressRequest(app, 'GET', '/integration/ai/stream', { + projectId, + eventId, + originalEventId, + }); + + expect(response.status).toBe(403); + expect(response.body.error).toBe('You have no access to this workspace'); + }); + + it('should return 400 when eventId is missing', async () => { + const app = setupApp(); + + const response = await makeExpressRequest(app, 'GET', '/integration/ai/stream', { + projectId, + originalEventId, + }); + + expect(response.status).toBe(400); + expect(response.body.error).toContain('eventId'); + }); + + it('should return 400 when originalEventId is missing', async () => { + const app = setupApp(); + + const response = await makeExpressRequest(app, 'GET', '/integration/ai/stream', { + projectId, + eventId, + }); + + expect(response.status).toBe(400); + expect(response.body.error).toContain('originalEventId'); + }); + + it('should return 404 when the event is not found', async () => { + mockStreamSuggestion.mockRejectedValue(new Error('Event not found')); + const app = setupApp(); + + const response = await makeExpressRequest(app, 'GET', '/integration/ai/stream', { + projectId, + eventId, + originalEventId, + }); + + expect(response.status).toBe(404); + expect(response.body.error).toBe('Event not found'); + }); + + it('should proxy the AI suggestion stream with the gateway status, headers and full SSE body', async () => { + mockStreamSuggestion.mockResolvedValue({ toUIMessageStreamResponse: () => createFakeStreamResponse() }); + const app = setupApp(); + + const response = await makeExpressRequest(app, 'GET', '/integration/ai/stream', { + projectId, + eventId, + originalEventId, + }); + + expect(mockGetEventsFactory).toHaveBeenCalledWith(expect.objectContaining({ user: expect.objectContaining({ id: userId }) }), projectId); + expect(mockStreamSuggestion).toHaveBeenCalledWith({}, eventId, originalEventId); + expect(response.status).toBe(200); + expect(response.headers['content-type']).toBe('text/event-stream'); + expect(response.body).toContain('data: {"type":"start"}'); + expect(response.body).toContain('data: {"type":"reasoning-delta","id":"reasoning-0","delta":"Reasoning"}'); + expect(response.body).toContain('data: [DONE]'); + }); +}); diff --git a/test/integrations/github-routes.test.ts b/test/integrations/github-routes.test.ts index 03eacc94a..1db61becd 100644 --- a/test/integrations/github-routes.test.ts +++ b/test/integrations/github-routes.test.ts @@ -3,6 +3,7 @@ import { ObjectId } from 'mongodb'; import express from 'express'; import { createGitHubRouter } from '../../src/integrations/github/routes'; import { ContextFactories } from '../../src/types/graphql'; +import { makeExpressRequest } from '../helpers/expressRequest'; /** * Mock GitHubService @@ -72,87 +73,6 @@ function createMockWorkspace(options: { }; } -/** - * Helper function to make a request to Express app - */ -function makeRequest( - app: express.Application, - method: string, - path: string, - query?: Record -): Promise<{ status: number; body: any }> { - return new Promise((resolve, reject) => { - const url = query ? `${path}?${new URLSearchParams(query).toString()}` : path; - const req = { - method, - url, - originalUrl: url, - path, - query: query || {}, - headers: {}, - get: jest.fn(), - params: {}, - body: {}, - } as any; - - let statusCode = 200; - let jsonCalled = false; - const res = { - status: (code: number) => { - statusCode = code; - - return res; - }, - json: (data: any) => { - jsonCalled = true; - resolve({ - status: statusCode, - body: data, - }); - }, - setHeader: jest.fn(), - getHeader: jest.fn(), - end: jest.fn(), - send: jest.fn((data?: any) => { - if (!jsonCalled) { - resolve({ - status: statusCode, - body: data, - }); - } - }), - redirect: jest.fn((redirectUrl: string) => { - statusCode = 302; - resolve({ - status: statusCode, - body: redirectUrl, - }); - }), - } as any; - - /** - * Use (app as any).handle() as handle method exists but is not in TypeScript types - * This simulates how Express processes requests internally - */ - (app as any).handle(req, res, (err: any) => { - if (err) { - reject(err); - } else if (!jsonCalled) { - /** - * If json was not called, check if response was sent another way - * Wait a bit to allow async handlers to complete - */ - setTimeout(() => { - resolve({ - status: statusCode, - body: null, - }); - }, 50); - } - }); - }); -} - describe('GitHub Routes - /integration/github/connect', () => { let app: express.Application; const userId = '507f1f77bcf86cd799439011'; @@ -242,7 +162,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/connect', { projectId }); + const response = await makeExpressRequest(app, 'GET', '/integration/github/connect', { projectId }); expect(response.status).toBe(200); expect(response.body).toHaveProperty('redirectUrl'); @@ -265,7 +185,7 @@ describe('GitHub Routes - /integration/github/connect', () => { req.context.user.id = undefined; }); - const response = await makeRequest(app, 'GET', '/integration/github/connect', { projectId }); + const response = await makeExpressRequest(app, 'GET', '/integration/github/connect', { projectId }); expect(response.status).toBe(401); expect(response.body).toHaveProperty('error'); @@ -284,7 +204,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/connect'); + const response = await makeExpressRequest(app, 'GET', '/integration/github/connect'); expect(response.status).toBe(400); expect(response.body).toHaveProperty('error'); @@ -303,7 +223,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/connect', { projectId: 'invalid-id' }); + const response = await makeExpressRequest(app, 'GET', '/integration/github/connect', { projectId: 'invalid-id' }); expect(response.status).toBe(400); expect(response.body).toHaveProperty('error'); @@ -325,7 +245,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/connect', { projectId }); + const response = await makeExpressRequest(app, 'GET', '/integration/github/connect', { projectId }); expect(response.status).toBe(404); expect(response.body).toHaveProperty('error'); @@ -351,7 +271,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/connect', { projectId }); + const response = await makeExpressRequest(app, 'GET', '/integration/github/connect', { projectId }); expect(response.status).toBe(400); expect(response.body).toHaveProperty('error'); @@ -385,7 +305,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/connect', { projectId }); + const response = await makeExpressRequest(app, 'GET', '/integration/github/connect', { projectId }); expect(response.status).toBe(403); expect(response.body).toHaveProperty('error'); @@ -419,7 +339,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/connect', { projectId }); + const response = await makeExpressRequest(app, 'GET', '/integration/github/connect', { projectId }); expect(response.status).toBe(403); expect(response.body).toHaveProperty('error'); @@ -474,7 +394,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { state, }); @@ -495,7 +415,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, }); @@ -518,7 +438,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, }); @@ -549,7 +469,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, }); @@ -582,7 +502,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, // eslint-disable-next-line @typescript-eslint/camelcase, camelcase @@ -619,7 +539,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, // eslint-disable-next-line @typescript-eslint/camelcase, camelcase @@ -657,7 +577,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, // eslint-disable-next-line @typescript-eslint/camelcase, camelcase @@ -697,7 +617,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, // eslint-disable-next-line @typescript-eslint/camelcase, camelcase @@ -732,7 +652,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, }); @@ -776,7 +696,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, }); @@ -837,7 +757,7 @@ describe('GitHub Routes - /integration/github/connect', () => { /** * OAuth callback without installation_id (installation already exists) */ - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, }); @@ -929,7 +849,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, // eslint-disable-next-line @typescript-eslint/camelcase, camelcase @@ -999,7 +919,7 @@ describe('GitHub Routes - /integration/github/connect', () => { setupRouter(factories); - const response = await makeRequest(app, 'GET', '/integration/github/oauth', { + const response = await makeExpressRequest(app, 'GET', '/integration/github/oauth', { code, state, // eslint-disable-next-line @typescript-eslint/camelcase, camelcase diff --git a/test/integrations/vercel-ai.test.ts b/test/integrations/vercel-ai.test.ts index a62347054..624b26d82 100644 --- a/test/integrations/vercel-ai.test.ts +++ b/test/integrations/vercel-ai.test.ts @@ -1,9 +1,10 @@ import '../../src/env-test'; -import { generateText } from 'ai'; +import { generateText, streamText } from 'ai'; import { vercelAIApi } from '../../src/integrations/vercel-ai/'; jest.mock('ai', () => ({ generateText: jest.fn(), + streamText: jest.fn(), })); describe('VercelAIApi', () => { @@ -38,4 +39,25 @@ describe('VercelAIApi', () => { expect(result).toBe('model output'); }); }); + + describe('stream', () => { + it('should forward the system/prompt pair to streamText and return its result synchronously', () => { + const streamResult = { toUIMessageStreamResponse: jest.fn() }; + + (streamText as jest.Mock).mockReturnValue(streamResult); + + const result = vercelAIApi.stream({ + system: testSystem, + prompt: testPrompt, + }); + + expect(streamText).toHaveBeenCalledWith({ + model: testModelId, + system: testSystem, + prompt: testPrompt, + providerOptions: testProviderOptions, + }); + expect(result).toBe(streamResult); + }); + }); }); diff --git a/test/services/askAi.test.ts b/test/services/askAi.test.ts index df056aa06..2c80cd8f6 100644 --- a/test/services/askAi.test.ts +++ b/test/services/askAi.test.ts @@ -10,6 +10,7 @@ import { SUGGESTION_FALLBACK_MESSAGE } from '../../src/services/askAi/security/l jest.mock('../../src/integrations/vercel-ai/', () => ({ vercelAIApi: { complete: jest.fn(), + stream: jest.fn(), }, })); @@ -21,7 +22,7 @@ jest.mock('@hawk.so/nodejs', () => ({ /** * Extract the per-request nonce from the prompt handed to the transport * - * @param prompt - prompt captured from the transport's `complete` call + * @param prompt - prompt captured from the transport's `complete`/`stream` call * @returns {string} nonce carried by the untrusted-data marker */ function nonceFromPrompt(prompt: string): string { @@ -127,4 +128,28 @@ describe('AIService', () => { expect(reported).not.toContain('Service marker'); }); }); + + describe('streamSuggestion', () => { + it('should spotlight the event with a nonce the system instruction repeats, and return the stream unchanged', async () => { + const streamResult = { toUIMessageStreamResponse: jest.fn() }; + + (vercelAIApi.stream as jest.Mock).mockReturnValue(streamResult); + + const result = await aiService.streamSuggestion(eventsFactoryWithPayload(), testEventId, testOriginalEventId); + const args = (vercelAIApi.stream as jest.Mock).mock.calls[0][0] as { system: string; prompt: string }; + + expect(args.prompt).toContain(JSON.stringify(testPayload)); + expect(args.system.startsWith(ctoInstruction)).toBe(true); + expect(args.system).toContain(nonceFromPrompt(args.prompt)); + expect(result).toBe(streamResult); + }); + + it('should throw Event not found when the events factory returns nothing', async () => { + await expect( + aiService.streamSuggestion(createEventsFactory(null), testEventId, testOriginalEventId) + ).rejects.toThrow('Event not found'); + + expect(vercelAIApi.stream).not.toHaveBeenCalled(); + }); + }); });