AI SDK v5 migration (#14549)

Co-authored-by: Félix Malfait <felix@twenty.com>
This commit is contained in:
Abdul Rahman
2025-09-23 01:43:43 +05:30
committed by GitHub
parent e8121919bd
commit 216d72b5d7
75 changed files with 1347 additions and 841 deletions
@@ -1,8 +1,12 @@
import { Test, type TestingModule } from '@nestjs/testing';
import { openai } from '@ai-sdk/openai';
import { ModelProvider } from 'src/engine/core-modules/ai/constants/ai-models.const';
import { AIBillingService } from 'src/engine/core-modules/ai/services/ai-billing.service';
import { AiModelRegistryService } from 'src/engine/core-modules/ai/services/ai-model-registry.service';
import { AiService } from 'src/engine/core-modules/ai/services/ai.service';
import { FeatureFlagService } from 'src/engine/core-modules/feature-flag/services/feature-flag.service';
import { AIBillingService } from 'src/engine/core-modules/ai/services/ai-billing.service';
import { AiController } from './ai.controller';
@@ -11,6 +15,7 @@ describe('AiController', () => {
let aiService: jest.Mocked<AiService>;
let featureFlagService: jest.Mocked<FeatureFlagService>;
let aiBillingService: jest.Mocked<AIBillingService>;
let aiModelRegistryService: jest.Mocked<AiModelRegistryService>;
beforeEach(async () => {
const mockAiService = {
@@ -26,6 +31,14 @@ describe('AiController', () => {
calculateAndBillUsage: jest.fn(),
};
const mockAiModelRegistryService = {
getDefaultModel: jest.fn().mockReturnValue({
modelId: 'gpt-4o',
provider: ModelProvider.OPENAI,
model: openai('gpt-4o'),
}),
};
const module: TestingModule = await Test.createTestingModule({
controllers: [AiController],
providers: [
@@ -41,6 +54,10 @@ describe('AiController', () => {
provide: AIBillingService,
useValue: mockAIBillingService,
},
{
provide: AiModelRegistryService,
useValue: mockAiModelRegistryService,
},
],
}).compile();
@@ -48,6 +65,7 @@ describe('AiController', () => {
aiService = module.get(AiService);
featureFlagService = module.get(FeatureFlagService);
aiBillingService = module.get(AIBillingService);
aiModelRegistryService = module.get(AiModelRegistryService);
});
it('should be defined', () => {
@@ -61,7 +79,7 @@ describe('AiController', () => {
const mockRequest = {
messages: [{ role: 'user' as const, content: 'Hello' }],
temperature: 0.7,
maxTokens: 100,
maxOutputTokens: 100,
};
const mockRes = {
@@ -70,19 +88,23 @@ describe('AiController', () => {
end: jest.fn(),
} as any;
const mockModel = { modelId: 'gpt-4o' } as any;
const mockModel = openai('gpt-4o');
aiService.getModel.mockReturnValue(mockModel);
aiModelRegistryService.getDefaultModel.mockReturnValue({
modelId: 'gpt-4o',
provider: ModelProvider.OPENAI,
model: mockModel,
});
const mockUsage = {
promptTokens: 10,
completionTokens: 20,
inputTokens: 10,
outputTokens: 20,
totalTokens: 30,
};
const mockStreamTextResult = {
usage: Promise.resolve(mockUsage),
pipeDataStreamToResponse: jest.fn(),
pipeUIMessageStreamToResponse: jest.fn(),
};
aiService.streamText.mockReturnValue(mockStreamTextResult as any);
@@ -96,12 +118,12 @@ describe('AiController', () => {
messages: mockRequest.messages,
options: {
temperature: 0.7,
maxTokens: 100,
maxOutputTokens: 100,
model: mockModel,
},
});
expect(
mockStreamTextResult.pipeDataStreamToResponse,
mockStreamTextResult.pipeUIMessageStreamToResponse,
).toHaveBeenCalledWith(mockRes);
expect(aiBillingService.calculateAndBillUsage).toHaveBeenCalledWith(
mockModel.modelId,
@@ -131,7 +153,11 @@ describe('AiController', () => {
const mockRes = {} as any;
aiService.getModel.mockReturnValue({ modelId: 'gpt-4o' } as any);
aiModelRegistryService.getDefaultModel.mockReturnValue({
modelId: 'gpt-4o',
provider: ModelProvider.OPENAI,
model: openai('gpt-4o'),
});
aiService.streamText.mockImplementation(() => {
throw new Error('Service error');
});
@@ -8,21 +8,22 @@ import {
UseGuards,
} from '@nestjs/common';
import { type CoreMessage } from 'ai';
import { type ModelMessage } from 'ai';
import { Response } from 'express';
import { AIBillingService } from 'src/engine/core-modules/ai/services/ai-billing.service';
import { AiModelRegistryService } from 'src/engine/core-modules/ai/services/ai-model-registry.service';
import { AiService } from 'src/engine/core-modules/ai/services/ai.service';
import { FeatureFlagKey } from 'src/engine/core-modules/feature-flag/enums/feature-flag-key.enum';
import { FeatureFlagService } from 'src/engine/core-modules/feature-flag/services/feature-flag.service';
import { Workspace } from 'src/engine/core-modules/workspace/workspace.entity';
import { AuthWorkspace } from 'src/engine/decorators/auth/auth-workspace.decorator';
import { WorkspaceAuthGuard } from 'src/engine/guards/workspace-auth.guard';
import { AIBillingService } from 'src/engine/core-modules/ai/services/ai-billing.service';
export interface ChatRequest {
messages: CoreMessage[];
messages: ModelMessage[];
temperature?: number;
maxTokens?: number;
maxOutputTokens?: number;
}
@Controller('chat')
@@ -32,6 +33,7 @@ export class AiController {
private readonly aiService: AiService,
private readonly featureFlagService: FeatureFlagService,
private readonly aiBillingService: AIBillingService,
private readonly aiModelRegistryService: AiModelRegistryService,
) {}
@Post()
@@ -52,7 +54,7 @@ export class AiController {
);
}
const { messages, temperature, maxTokens } = request;
const { messages, temperature, maxOutputTokens } = request;
if (!messages || messages.length === 0) {
throw new HttpException(
@@ -62,27 +64,26 @@ export class AiController {
}
try {
// TODO: Add support for custom models
const model = this.aiService.getModel(undefined);
const registeredModel = this.aiModelRegistryService.getDefaultModel();
const result = this.aiService.streamText({
messages,
options: {
temperature,
maxTokens,
model,
maxOutputTokens,
model: registeredModel.model,
},
});
result.usage.then((usage) => {
this.aiBillingService.calculateAndBillUsage(
model.modelId,
registeredModel.modelId,
usage,
workspace.id,
);
});
result.pipeDataStreamToResponse(res);
result.pipeUIMessageStreamToResponse(res);
} catch (error) {
const errorMessage =
error instanceof Error ? error.message : 'Unknown error occurred';
@@ -11,8 +11,8 @@ describe('AIBillingService', () => {
let mockWorkspaceEventEmitter: jest.Mocked<WorkspaceEventEmitter>;
const mockTokenUsage = {
promptTokens: 1000,
completionTokens: 500,
inputTokens: 1000,
outputTokens: 500,
totalTokens: 1500,
};
@@ -63,8 +63,8 @@ describe('AIBillingService', () => {
it('should calculate cost correctly with different token usage', async () => {
const differentTokenUsage = {
promptTokens: 2000,
completionTokens: 1000,
inputTokens: 2000,
outputTokens: 1000,
totalTokens: 3000,
};
@@ -1,17 +1,19 @@
import { Test, type TestingModule } from '@nestjs/testing';
import { HttpException, HttpStatus } from '@nestjs/common';
import { Test, type TestingModule } from '@nestjs/testing';
import { getRepositoryToken } from '@nestjs/typeorm';
import { jsonSchema } from 'ai';
import { MCP_SERVER_METADATA } from 'src/engine/core-modules/ai/constants/mcp.const';
import { type JsonRpc } from 'src/engine/core-modules/ai/dtos/json-rpc';
import { McpService } from 'src/engine/core-modules/ai/services/mcp.service';
import { ToolService } from 'src/engine/core-modules/ai/services/tool.service';
import { FeatureFlagKey } from 'src/engine/core-modules/feature-flag/enums/feature-flag-key.enum';
import { FeatureFlagService } from 'src/engine/core-modules/feature-flag/services/feature-flag.service';
import { UserRoleService } from 'src/engine/metadata-modules/user-role/user-role.service';
import { ToolService } from 'src/engine/core-modules/ai/services/tool.service';
import { type Workspace } from 'src/engine/core-modules/workspace/workspace.entity';
import { type JsonRpc } from 'src/engine/core-modules/ai/dtos/json-rpc';
import { MCP_SERVER_METADATA } from 'src/engine/core-modules/ai/constants/mcp.const';
import { ADMIN_ROLE_LABEL } from 'src/engine/metadata-modules/permissions/constants/admin-role-label.constants';
import { RoleEntity } from 'src/engine/metadata-modules/role/role.entity';
import { McpService } from 'src/engine/core-modules/ai/services/mcp.service';
import { UserRoleService } from 'src/engine/metadata-modules/user-role/user-role.service';
describe('McpService', () => {
let service: McpService;
@@ -197,7 +199,7 @@ describe('McpService', () => {
const mockTool = {
description: 'Test tool',
parameters: { jsonSchema: { type: 'object', properties: {} } },
inputSchema: jsonSchema({ type: 'object', properties: {} }),
execute: jest.fn().mockResolvedValue({ result: 'success' }),
};
@@ -245,7 +247,7 @@ describe('McpService', () => {
const mockTool = {
description: 'Test tool',
parameters: { jsonSchema: { type: 'object', properties: {} } },
inputSchema: jsonSchema({ type: 'object', properties: {} }),
execute: jest.fn().mockResolvedValue({ result: 'success' }),
};
@@ -299,7 +301,7 @@ describe('McpService', () => {
const mockToolsMap = {
testTool: {
description: 'Test tool',
parameters: { jsonSchema: { type: 'object', properties: {} } },
inputSchema: jsonSchema({ type: 'object', properties: {} }),
},
};
@@ -328,7 +330,7 @@ describe('McpService', () => {
{
name: 'testTool',
description: 'Test tool',
inputSchema: { type: 'object', properties: {} },
inputSchema: jsonSchema({ type: 'object', properties: {} }),
},
],
}),
@@ -1,5 +1,7 @@
import { Test } from '@nestjs/testing';
import { jsonSchema } from 'ai';
import { ToolAdapterService } from 'src/engine/core-modules/ai/services/tool-adapter.service';
import { ToolType } from 'src/engine/core-modules/tool/enums/tool-type.enum';
import { ToolRegistryService } from 'src/engine/core-modules/tool/services/tool-registry.service';
@@ -33,7 +35,7 @@ describe('ToolAdapterService', () => {
}));
const unflaggedTool: Tool = {
description: 'HTTP Request tool',
parameters: { type: 'object', properties: {} },
inputSchema: jsonSchema({ type: 'object', properties: {} }),
execute: unflaggedToolExecute,
};
@@ -44,7 +46,7 @@ describe('ToolAdapterService', () => {
}));
const flaggedTool: Tool = {
description: 'Send Email tool',
parameters: { type: 'object', properties: {} },
inputSchema: jsonSchema({ type: 'object', properties: {} }),
execute: flaggedToolExecute,
flag: PermissionFlagType.SEND_EMAIL_TOOL,
};
@@ -1,5 +1,7 @@
import { Injectable, Logger } from '@nestjs/common';
import { LanguageModelUsage } from 'ai';
import { type ModelId } from 'src/engine/core-modules/ai/constants/ai-models.const';
import { DOLLAR_TO_CREDIT_MULTIPLIER } from 'src/engine/core-modules/ai/constants/dollar-to-credit-multiplier';
import { AiModelRegistryService } from 'src/engine/core-modules/ai/services/ai-model-registry.service';
@@ -8,12 +10,6 @@ import { BillingMeterEventName } from 'src/engine/core-modules/billing/enums/bil
import { type BillingUsageEvent } from 'src/engine/core-modules/billing/types/billing-usage-event.type';
import { WorkspaceEventEmitter } from 'src/engine/workspace-event-emitter/workspace-event-emitter';
export interface TokenUsage {
promptTokens: number;
completionTokens: number;
totalTokens: number;
}
@Injectable()
export class AIBillingService {
private readonly logger = new Logger(AIBillingService.name);
@@ -23,7 +19,10 @@ export class AIBillingService {
private readonly aiModelRegistryService: AiModelRegistryService,
) {}
async calculateCost(modelId: ModelId, usage: TokenUsage): Promise<number> {
async calculateCost(
modelId: ModelId,
usage: LanguageModelUsage,
): Promise<number> {
const model = this.aiModelRegistryService.getEffectiveModelConfig(modelId);
if (!model) {
@@ -31,9 +30,9 @@ export class AIBillingService {
}
const inputCost =
(usage.promptTokens / 1000) * model.inputCostPer1kTokensInCents;
((usage.inputTokens ?? 0) / 1000) * model.inputCostPer1kTokensInCents;
const outputCost =
(usage.completionTokens / 1000) * model.outputCostPer1kTokensInCents;
((usage.outputTokens ?? 0) / 1000) * model.outputCostPer1kTokensInCents;
const totalCost = inputCost + outputCost;
@@ -46,7 +45,7 @@ export class AIBillingService {
async calculateAndBillUsage(
modelId: ModelId,
usage: TokenUsage,
usage: LanguageModelUsage,
workspaceId: string,
): Promise<void> {
const costInCents = await this.calculateCost(modelId, usage);
@@ -138,7 +138,7 @@ export class AiModelRegistryService {
return Array.from(this.modelRegistry.values());
}
getDefaultModel(): RegisteredAIModel | undefined {
getDefaultModel(): RegisteredAIModel {
const defaultModelId = this.twentyConfigService.get('DEFAULT_MODEL_ID');
let model = this.getModel(defaultModelId);
@@ -1,6 +1,6 @@
import { Injectable } from '@nestjs/common';
import { type CoreMessage, streamText, LanguageModelV1 } from 'ai';
import { LanguageModel, type ModelMessage, streamText } from 'ai';
import { AiModelRegistryService } from 'src/engine/core-modules/ai/services/ai-model-registry.service';
@@ -28,18 +28,18 @@ export class AiService {
messages,
options,
}: {
messages: CoreMessage[];
messages: ModelMessage[];
options: {
temperature?: number;
maxTokens?: number;
model: LanguageModelV1;
maxOutputTokens?: number;
model: LanguageModel;
};
}) {
return streamText({
model: options.model,
messages,
temperature: options?.temperature,
maxTokens: options?.maxTokens,
maxOutputTokens: options?.maxOutputTokens,
});
}
}
@@ -204,11 +204,11 @@ export class McpService {
private handleToolsListing(id: string | number, toolSet: ToolSet) {
const toolsArray = Object.entries(toolSet)
.filter(([, def]) => !!def.parameters.jsonSchema)
.filter(([, def]) => !!def.inputSchema)
.map(([name, def]) => ({
name,
description: def.description,
inputSchema: def.parameters.jsonSchema,
inputSchema: def.inputSchema,
}));
return wrapJsonRpcResponse(id, {
@@ -42,7 +42,7 @@ export class ToolAdapterService {
private createToolSet(tool: Tool) {
return {
description: tool.description,
parameters: tool.parameters,
inputSchema: tool.inputSchema,
execute: async (parameters: { input: ToolInput }) =>
tool.execute(parameters.input),
};
@@ -61,7 +61,7 @@ export class ToolService {
if (objectPermission.canUpdate) {
tools[`create_${objectMetadata.nameSingular}`] = {
description: `Create a new ${objectMetadata.labelSingular} record. Provide all required fields and any optional fields you want to set. The system will automatically handle timestamps and IDs. Returns the created record with all its data.`,
parameters: getRecordInputSchema(objectMetadata),
inputSchema: getRecordInputSchema(objectMetadata),
execute: async (parameters) => {
return this.createRecord(
objectMetadata.nameSingular,
@@ -74,7 +74,7 @@ export class ToolService {
tools[`update_${objectMetadata.nameSingular}`] = {
description: `Update an existing ${objectMetadata.labelSingular} record. Provide the record ID and only the fields you want to change. Unspecified fields will remain unchanged. Returns the updated record with all current data.`,
parameters: getRecordInputSchema(objectMetadata),
inputSchema: getRecordInputSchema(objectMetadata),
execute: async (parameters) => {
return this.updateRecord(
objectMetadata.nameSingular,
@@ -89,7 +89,7 @@ export class ToolService {
if (objectPermission.canRead) {
tools[`find_${objectMetadata.nameSingular}`] = {
description: `Search for ${objectMetadata.labelSingular} records using flexible filtering criteria. Supports exact matches, pattern matching, ranges, and null checks. Use limit/offset for pagination. Returns an array of matching records with their full data.`,
parameters: generateFindToolSchema(objectMetadata),
inputSchema: generateFindToolSchema(objectMetadata),
execute: async (parameters) => {
return this.findRecords(
objectMetadata.nameSingular,
@@ -102,7 +102,7 @@ export class ToolService {
tools[`find_one_${objectMetadata.nameSingular}`] = {
description: `Retrieve a single ${objectMetadata.labelSingular} record by its unique ID. Use this when you know the exact record ID and need the complete record data. Returns the full record or an error if not found.`,
parameters: generateFindOneToolSchema(),
inputSchema: generateFindOneToolSchema(),
execute: async (parameters) => {
return this.findOneRecord(
objectMetadata.nameSingular,
@@ -117,7 +117,7 @@ export class ToolService {
if (objectPermission.canSoftDelete) {
tools[`soft_delete_${objectMetadata.nameSingular}`] = {
description: `Soft delete a ${objectMetadata.labelSingular} record by marking it as deleted. The record remains in the database but is hidden from normal queries. This is reversible and preserves all data. Use this for temporary removal.`,
parameters: generateSoftDeleteToolSchema(),
inputSchema: generateSoftDeleteToolSchema(),
execute: async (parameters) => {
return this.softDeleteRecord(
objectMetadata.nameSingular,
@@ -130,7 +130,7 @@ export class ToolService {
tools[`soft_delete_many_${objectMetadata.nameSingular}`] = {
description: `Soft delete multiple ${objectMetadata.labelSingular} records at once by providing an array of record IDs. All records are marked as deleted but remain in the database. This is efficient for bulk operations and preserves all data.`,
parameters: generateBulkDeleteToolSchema(),
inputSchema: generateBulkDeleteToolSchema(),
execute: async (parameters) => {
return this.softDeleteManyRecords(
objectMetadata.nameSingular,