From 0a7df48daad64323725e37da463fc2cf23ca6ef2 Mon Sep 17 00:00:00 2001 From: Reversean Date: Wed, 29 Jul 2026 19:34:51 +0300 Subject: [PATCH] feat(ai): Ask AI suggestions streaming via HTTP New HTTP route GET /integration/ai/stream added. This route calls Ask AI service about specified event and responds text/event-stream. Reponse carrying AiStream which represent AiStreamPart sequence: either text-delta (text-increments generated by AI assistant) or error (failure description during response generation). NOTE: Response doesn't carry reasoning, tooling and start/end parts since they're not required yet. Route checks workspace membership before calling Ask AI. For this purpose function checkUserInWorkspaceByProjectId became exported. Failed membership check leads to response with 403. Also route checks if specified event exists. Failed check leads to response with 404. Aborting request cancel Ask AI suggestion generation. For this purpose AbortController is declared as an eslint global: it is on globalThis since Node 15, but eslint's node env predates it. --- .eslintrc.js | 6 + package.json | 4 +- src/directives/requireUserInWorkspace.ts | 2 +- src/index.ts | 6 + src/integrations/vercel-ai/index.ts | 79 ++++++- src/services/askAi/index.ts | 1 + src/services/askAi/routes.ts | 156 +++++++++++++ src/services/askAi/service.ts | 73 ++++++- src/services/types.ts | 4 +- test/helpers/expressRequest.ts | 173 +++++++++++++++ test/integrations/github-routes.test.ts | 124 ++--------- test/integrations/vercel-ai.test.ts | 76 ++++++- test/services/askAi.test.ts | 44 +++- test/services/askAiRoutes.test.ts | 265 +++++++++++++++++++++++ yarn.lock | 8 +- 15 files changed, 895 insertions(+), 126 deletions(-) create mode 100644 src/services/askAi/routes.ts create mode 100644 test/helpers/expressRequest.ts create mode 100644 test/services/askAiRoutes.test.ts diff --git a/.eslintrc.js b/.eslintrc.js index 12245e12..29300a3d 100644 --- a/.eslintrc.js +++ b/.eslintrc.js @@ -4,6 +4,12 @@ module.exports = { 'node': true, 'jest': true }, + globals: { + /** + * TODO: bump eslint since it's current env uses older "node" version which missing required global types + */ + 'AbortController': 'readonly' + }, rules: { '@typescript-eslint/camelcase': 'warn', '@typescript-eslint/no-unused-vars': 'warn', diff --git a/package.json b/package.json index db9cff2d..d0757d5f 100644 --- a/package.json +++ b/package.json @@ -1,6 +1,6 @@ { "name": "hawk.api", - "version": "1.5.13", + "version": "1.5.14", "main": "index.ts", "license": "BUSL-1.1", "scripts": { @@ -42,7 +42,7 @@ "@graphql-tools/schema": "^8.5.1", "@graphql-tools/utils": "^8.9.0", "@hawk.so/nodejs": "^3.3.2", - "@hawk.so/types": "^0.5.9", + "@hawk.so/types": "^0.7.0", "@n1ru4l/json-patch-plus": "^0.2.0", "@node-saml/node-saml": "^5.0.1", "@octokit/oauth-methods": "^4.0.0", 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..c60d2bfe 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 './services/askAi'; /** * 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 6dfe0e83..d931614b 100644 --- a/src/integrations/vercel-ai/index.ts +++ b/src/integrations/vercel-ai/index.ts @@ -1,4 +1,6 @@ -import { generateText } from 'ai'; +import { generateText, streamText, type TextStreamPart, type ToolSet } from 'ai'; +import { getErrorMessage, ProviderOptions } from '@ai-sdk/provider-utils'; +import type { AiStream } from '@hawk.so/types'; /** * Params for a single completion call to the model @@ -15,6 +17,44 @@ export interface CompletionParams { prompt: string; } +/** + * Params for a streaming completion call to the model + */ +export interface StreamParams extends CompletionParams { + /** + * Aborted when the answer is no longer required, which stops the model + */ + signal: AbortSignal; +} + +/** + * Converts Vercel SDK's stream parts. + * + * Everything but text and error parts is dropped. + * + * @param parts - stream of incoming SDK parts + * @returns {AiStream} stream converted of converted parts + */ +async function * toAiStream( + parts: AsyncIterable> +): AiStream { + for await (const part of parts) { + if (part.type === 'text-delta') { + yield { + type: 'text-delta', + delta: part.text, + }; + } + + if (part.type === 'error') { + yield { + type: 'error', + errorText: getErrorMessage(part.error), + }; + } + } +} + /** * Interface for interacting with Vercel AI Gateway * @@ -27,11 +67,24 @@ class VercelAIApi { */ private readonly modelId: string; + /** + * Provider Gateway configurations + */ + 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'], + }, + }; } /** @@ -45,15 +98,29 @@ 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 streamed text + * + * @param {StreamParams} params - system instruction, prompt and abort signal + * @returns {AiStream} text generated by the model, as it arrives + */ + public stream({ system, prompt, signal }: StreamParams): AiStream { + const { fullStream } = streamText({ + model: this.modelId, + system, + prompt, + providerOptions: this.providerOptions, + abortSignal: signal, + }); + + return toAiStream(fullStream); + } } export const vercelAIApi = new VercelAIApi(); diff --git a/src/services/askAi/index.ts b/src/services/askAi/index.ts index 5a5e2b63..f903e022 100644 --- a/src/services/askAi/index.ts +++ b/src/services/askAi/index.ts @@ -1 +1,2 @@ export { AskAiService, askAiService } from './service'; +export { appendAiAssistantRoutes } from './routes'; diff --git a/src/services/askAi/routes.ts b/src/services/askAi/routes.ts new file mode 100644 index 00000000..fef8cc88 --- /dev/null +++ b/src/services/askAi/routes.ts @@ -0,0 +1,156 @@ +import '../../typeDefs/expressContext'; +import express from 'express'; +import { ObjectId } from 'mongodb'; +import { getEventsFactory } from '../../resolvers/helpers/eventsFactory'; +import { checkUserInWorkspaceByProjectId } from '../../directives/requireUserInWorkspace'; +import { askAiService } from './service'; +import type { AiStreamPart } from '@hawk.so/types'; + +/** + * 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 (may be `string[]` if repeated) + * @returns user id and validated project id if authorized, `null` otherwise (response already sent) + */ +async function authorizeProjectAccess( + req: express.Request, + res: express.Response, + projectId: unknown +): Promise<{ userId: string; projectId: string } | null> { + const userId = req.context?.user?.id; + + if (!userId) { + res.status(401).json({ error: 'Unauthorized. Please provide authorization token.' }); + + return null; + } + + if (!projectId || typeof projectId !== 'string') { + res.status(400).json({ error: 'projectId query parameter is required' }); + + return null; + } + + if (!ObjectId.isValid(projectId)) { + res.status(400).json({ error: `Invalid projectId format: ${projectId}` }); + + 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, + projectId, + }; +} + +/** + * 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) => { + const abort = new AbortController(); + + /** Abort response generation when connection is closed */ + res.on('close', () => abort.abort()); + + try { + const { projectId, eventId, originalEventId } = req.query; + + const authResult = await authorizeProjectAccess(req, res, projectId); + + if (!authResult) { + 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, authResult.projectId); + + let stream; + + try { + stream = await askAiService.streamSuggestion(eventsFactory, eventId, originalEventId, abort.signal); + } catch (error) { + if (!(error instanceof Error) || error.message !== 'Event not found') { + throw error; + } + + res.status(404).json({ error: error.message }); + + return; + } + + res.writeHead(200, { + 'content-type': 'text/event-stream', + 'cache-control': 'no-cache', + connection: 'keep-alive', + }); + + try { + for await (const part of stream) { + if (abort.signal.aborted) { + break; + } + + res.write(`data: ${JSON.stringify(part)}\n\n`); + } + } catch (error) { + if (!abort.signal.aborted) { + const part: AiStreamPart = { + type: 'error', + errorText: error instanceof Error ? error.message : 'AI suggestion failed.', + }; + + res.write(`data: ${JSON.stringify(part)}\n\n`); + } + } + + res.end(); + } catch (error) { + if (abort.signal.aborted) { + return; + } + + 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/askAi/service.ts b/src/services/askAi/service.ts index f930af75..b38be3a3 100644 --- a/src/services/askAi/service.ts +++ b/src/services/askAi/service.ts @@ -4,6 +4,8 @@ import { buildEventPrompt, spotlightInstruction } from './security/spotlighting' import { echoesNonce, SUGGESTION_FALLBACK_MESSAGE } from './security/nonceEcho'; import { ctoInstruction } from './instructions/cto'; import { EventsFactoryInterface } from '../types'; +import type { Event } from '../types'; +import type { AiStream } from '@hawk.so/types'; /** * Report that the nonce check rejected an answer. @@ -26,7 +28,8 @@ function reportRejectedSuggestion(eventId: string, originalEventId: string): voi } /** - * Service for interacting with AI + * Looks up an event and turns it into an AI suggestion, guarding the model + * call against the event payload trying to hijack the prompt. */ export class AskAiService { /** @@ -40,12 +43,12 @@ export class AskAiService { * @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); @@ -62,6 +65,62 @@ export class AskAiService { return text; } + + /** + * Generate a streaming suggestion for the event, spotlighted exactly as + * {@link AskAiService.generateSuggestion} + * + * @param eventsFactory - events factory + * @param eventId - event id + * @param originalEventId - original event id + * @param signal - aborted when the answer is no longer wanted + * @returns {Promise} - suggestion, as the model writes it + */ + public async streamSuggestion( + eventsFactory: EventsFactoryInterface, + eventId: string, + originalEventId: string, + signal: AbortSignal + ): Promise { + const event = await this.getEventOrThrow(eventsFactory, eventId, originalEventId); + + const { prompt, nonce } = buildEventPrompt(event.payload); + + return vercelAIApi.stream({ + system: ctoInstruction + spotlightInstruction(nonce), + prompt, + signal, + }); + } + + /** + * Find the event repetition. A failed lookup is reported as a missing one, + * so the reason does not reach the caller. + * + * @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 { + let event: Event | null; + + try { + event = await eventsFactory.getEventRepetition(eventId, originalEventId); + } catch { + throw new Error('Event not found'); + } + + if (!event) { + throw new Error('Event not found'); + } + + return event; + } } export const askAiService = new AskAiService(); 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..178c1ad6 --- /dev/null +++ b/test/helpers/expressRequest.ts @@ -0,0 +1,173 @@ +import { Writable } from 'stream'; +import express from 'express'; + +export interface CapturedResponse { + status: number; + headers: Record; + body: any; +} + +/** + * The part of Express's response API the routes under test use + */ +interface FakeResponse extends Writable { + status(code: number): FakeResponse; + setHeader(key: string, value: string): FakeResponse; + getHeader(key: string): string | undefined; + writeHead(statusCode: number, headers?: Record): FakeResponse; + json(data: any): void; + send(data?: any): void; + redirect(url: string): void; +} + +/** + * Rebind inherited methods as own properties, which a prototype swap cannot hide. + * + * Express's expressInit runs `setPrototypeOf(res, app.response)` on every request, + * which would otherwise leave the fake response with Node's real implementation + * reaching for a socket that does not exist here. + * + * @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 that records what a route wrote to it + * + * @param settle - called once with everything the route wrote to the response + * @returns {FakeResponse} fake response object to hand to Express + */ +function createFakeResponse(settle: (result: CapturedResponse) => void): FakeResponse { + 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 = 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(); + }, + }) as FakeResponse; + + pinInheritedMethodsAsOwnProperties(res); + + res.status = (code: number): FakeResponse => { + statusCode = code; + + return res; + }; + res.setHeader = (key: string, value: string): FakeResponse => { + headers[key] = value; + + return res; + }; + res.getHeader = (key: string): string | undefined => headers[key]; + res.writeHead = (statusCode_: number, newHeaders?: Record): FakeResponse => { + statusCode = statusCode_; + Object.assign(headers, newHeaders); + + return res; + }; + 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; +} + +/** + * Send a request through an Express app without opening a socket + * + * @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; an array value repeats the key + * @param onResponse - called with the response before the request is routed, for a test + * that has to act on it mid-flight + * @returns {Promise} status, headers and body the route produced + */ +export function makeExpressRequest( + app: express.Application, + method: string, + path: string, + query?: Record, + onResponse?: (res: FakeResponse) => void +): Promise { + return new Promise((resolve, reject) => { + const searchParams = new URLSearchParams(); + + for (const [key, value] of Object.entries(query || {})) { + for (const entry of Array.isArray(value) ? value : [ value ]) { + searchParams.append(key, entry); + } + } + + const url = query ? `${path}?${searchParams.toString()}` : path; + const req = { + method, + url, + originalUrl: url, + path, + query: query || {}, + headers: {}, + get: jest.fn(), + params: {}, + body: {}, + } as any; + + const res = createFakeResponse(resolve); + + if (onResponse) { + onResponse(res); + } + + (app as any).handle(req, res, (err: any) => { + if (err) { + reject(err); + } + }); + }); +} 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..d7e8784d 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,77 @@ describe('VercelAIApi', () => { expect(result).toBe('model output'); }); }); + + describe('stream', () => { + const testSignal = new AbortController().signal; + + /** + * Answer streamText with a canned stream of parts in the SDK's own shape + * + * @param parts - parts the model is to produce + */ + function modelProduces(parts: unknown[]): void { + (streamText as jest.Mock).mockReturnValue({ + fullStream: (async function * () { + yield* parts; + })(), + }); + } + + /** + * Read everything the adapter yields for the test prompt + * + * @returns {Promise} suggestion parts, in order + */ + async function readSuggestion(): Promise { + const parts = []; + + for await (const part of vercelAIApi.stream({ + system: testSystem, + prompt: testPrompt, + signal: testSignal, + })) { + parts.push(part); + } + + return parts; + } + + it('should forward the system/prompt pair and the abort signal to streamText', async () => { + modelProduces([]); + + await readSuggestion(); + + expect(streamText).toHaveBeenCalledWith({ + model: testModelId, + system: testSystem, + prompt: testPrompt, + providerOptions: testProviderOptions, + abortSignal: testSignal, + }); + }); + + it('should turn the model text deltas into text parts', async () => { + modelProduces([ + { type: 'start' }, + { type: 'reasoning-delta', id: '0', text: 'thinking out loud', }, + { type: 'text-delta', id: '0', text: 'Answer ', }, + { type: 'text-delta', id: '0', text: 'continues', }, + { type: 'finish' }, + ]); + + await expect(readSuggestion()).resolves.toEqual([ + { type: 'text-delta', delta: 'Answer ', }, + { type: 'text-delta', delta: 'continues', }, + ]); + }); + + it('should turn a model failure into an error part carrying its message', async () => { + modelProduces([{ type: 'error', error: new Error('gateway unavailable') }]); + + await expect(readSuggestion()).resolves.toEqual([ + { type: 'error', errorText: 'gateway unavailable', }, + ]); + }); + }); }); diff --git a/test/services/askAi.test.ts b/test/services/askAi.test.ts index 58e8a8e5..88a06f9b 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/n 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 { @@ -99,6 +100,18 @@ describe('AskAiService', () => { expect(vercelAIApi.complete).not.toHaveBeenCalled(); }); + it('should normalize a thrown lookup failure to Event not found', async () => { + const eventsFactory = { + getEventRepetition: jest.fn().mockRejectedValue(new Error(`Cant find event repetition for repetitionId: ${testEventId}`)), + }; + + await expect( + askAiService.generateSuggestion(eventsFactory, testEventId, testOriginalEventId) + ).rejects.toThrow('Event not found'); + + expect(vercelAIApi.complete).not.toHaveBeenCalled(); + }); + it('should return the fallback and report the event ids when the answer echoes the nonce', async () => { respondWithNonce(); @@ -128,4 +141,33 @@ describe('AskAiService', () => { expect(JSON.stringify(context)).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 = (async function * () { + yield { type: 'text-delta', delta: 'Answer' }; + })(); + + (vercelAIApi.stream as jest.Mock).mockReturnValue(streamResult); + + const signal = new AbortController().signal; + + const result = await askAiService.streamSuggestion(eventsFactoryWithPayload(), testEventId, testOriginalEventId, signal); + const args = (vercelAIApi.stream as jest.Mock).mock.calls[0][0] as { system: string; prompt: string; signal: AbortSignal }; + + expect(args.prompt).toContain(JSON.stringify(testPayload)); + expect(args.system.startsWith(ctoInstruction)).toBe(true); + expect(args.system).toContain(nonceFromPrompt(args.prompt)); + expect(args.signal).toBe(signal); + expect(result).toBe(streamResult); + }); + + it('should throw Event not found when the events factory returns nothing', async () => { + await expect( + askAiService.streamSuggestion(createEventsFactory(null), testEventId, testOriginalEventId, new AbortController().signal) + ).rejects.toThrow('Event not found'); + + expect(vercelAIApi.stream).not.toHaveBeenCalled(); + }); + }); }); diff --git a/test/services/askAiRoutes.test.ts b/test/services/askAiRoutes.test.ts new file mode 100644 index 00000000..b3dac66e --- /dev/null +++ b/test/services/askAiRoutes.test.ts @@ -0,0 +1,265 @@ +import '../../src/env-test'; +import express from 'express'; +import { makeExpressRequest } from '../helpers/expressRequest'; + +import { askAiService } from '../../src/services/askAi/service'; +import { getEventsFactory } from '../../src/resolvers/helpers/eventsFactory'; +import { checkUserInWorkspaceByProjectId } from '../../src/directives/requireUserInWorkspace'; +import { createAiStreamRouter, appendAiAssistantRoutes } from '../../src/services/askAi/routes'; +import type { AiStreamPart } from '@hawk.so/types'; + +jest.mock('../../src/services/askAi/service', () => ({ + askAiService: { + streamSuggestion: jest.fn(), + }, +})); + +jest.mock('../../src/resolvers/helpers/eventsFactory', () => ({ + getEventsFactory: jest.fn(), +})); + +jest.mock('../../src/directives/requireUserInWorkspace', () => ({ + checkUserInWorkspaceByProjectId: jest.fn(), +})); + +const mockStreamSuggestion = askAiService.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; +} + +/** + * Answer the route with a canned suggestion stream + * + * @param parts - parts the service hands to the route + */ +function aiStreamOf(parts: AiStreamPart[]): void { + mockStreamSuggestion.mockResolvedValue((async function * () { + yield* parts; + })()); +} + +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 401 when the request has no context at all', async () => { + const app = express(); + + app.use('/integration/ai', createAiStreamRouter()); + + 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 400 when projectId is repeated (parsed as an array)', async () => { + const app = setupApp(); + + const response = await makeExpressRequest(app, 'GET', '/integration/ai/stream', { + projectId: [projectId, 'another-project'], + eventId, + originalEventId, + }); + + expect(response.status).toBe(400); + expect(response.body.error).toContain('projectId'); + expect(mockCheckUserInWorkspaceByProjectId).not.toHaveBeenCalled(); + }); + + 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 generate the suggestion for the requested event of the authorized project', async () => { + aiStreamOf([ + { + type: 'text-delta', + delta: 'The stack trace ', + }, + { + type: 'text-delta', + delta: 'points at a null dereference', + }, + ]); + const app = setupApp(); + + const response = await makeExpressRequest(app, 'GET', '/integration/ai/stream', { + projectId, + eventId, + originalEventId, + }); + + expect(response.status).toBe(200); + expect(response.headers['content-type']).toBe('text/event-stream'); + expect(response.body).toBe( + 'data: {"type":"text-delta","delta":"The stack trace "}\n\n' + + 'data: {"type":"text-delta","delta":"points at a null dereference"}\n\n' + ); + }); + + it('should abort answer generation once the connection is closed', async () => { + let closeConnection = (): void => {}; + + mockStreamSuggestion.mockResolvedValue((async function * () { + yield { + type: 'text-delta', + delta: 'read by the client', + }; + closeConnection(); + yield { + type: 'text-delta', + delta: 'written after the client left', + }; + })()); + const app = setupApp(); + + const response = await makeExpressRequest( + app, + 'GET', + '/integration/ai/stream', + { + projectId, + eventId, + originalEventId, + }, + (res) => { + closeConnection = (): void => { + res.emit('close'); + }; + } + ); + + expect(response.body).toBe('data: {"type":"text-delta","delta":"read by the client"}\n\n'); + }); + + it('should send a error part when the underlying stream throws', async () => { + mockStreamSuggestion.mockResolvedValue((async function * () { + yield { + type: 'text-delta', + delta: 'The stack trace ', + }; + throw new Error('gateway unavailable'); + })()); + const app = setupApp(); + + const response = await makeExpressRequest(app, 'GET', '/integration/ai/stream', { + projectId, + eventId, + originalEventId, + }); + + expect(response.status).toBe(200); + expect(response.body).toBe( + 'data: {"type":"text-delta","delta":"The stack trace "}\n\n' + + 'data: {"type":"error","errorText":"gateway unavailable"}\n\n' + ); + }); +}); diff --git a/yarn.lock b/yarn.lock index ff56e295..7cb0a407 100644 --- a/yarn.lock +++ b/yarn.lock @@ -510,10 +510,10 @@ dependencies: bson "^7.0.0" -"@hawk.so/types@^0.5.9": - version "0.5.9" - resolved "https://registry.yarnpkg.com/@hawk.so/types/-/types-0.5.9.tgz#817e8b26283d0367371125f055f2e37a274797bc" - integrity sha512-86aE0Bdzvy8C+Dqd1iZpnDho44zLGX/t92SGuAv2Q52gjSJ7SHQdpGDWtM91FXncfT5uzAizl9jYMuE6Qrtm0Q== +"@hawk.so/types@^0.7.0": + version "0.7.0" + resolved "https://registry.yarnpkg.com/@hawk.so/types/-/types-0.7.0.tgz#ee959bc3d3ffa46c4c9d88693ea4d7f26c834f41" + integrity sha512-V8zCbnxwu1vVveZvfrm1Xoipi+hP5PZff/QiWeKGCl5/JeC0LX6AeCSIpKo2VvSDwiGcNTwmHygH6Tfch0Y2Dw== dependencies: bson "^7.0.0"