diff --git a/.eslintrc.js b/.eslintrc.js index 12245e12..7cec6e3f 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 a310be75..638a7e30 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 092b651b..1626cccd 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 cb6f8d93..897d2c94 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 745e2624..88812e7e 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 00000000..454f4f84 --- /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 5924cb5f..ed9d689b 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 1b14501f..2007767b 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 00000000..7e31bbfe --- /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 00000000..5f584bce --- /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 03eacc94..1db61bec 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 a6234705..624b26d8 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 df056aa0..2c80cd8f 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(); + }); + }); });