import { ModelMessage } from '@ai-sdk/provider-utils'; import { Inject } from '@nestjs/common'; import { LanguageModel, ToolSet, UIMessage, convertToModelMessages, createUIMessageStream, streamText } from 'ai'; import { z } from 'zod'; import { Input, RunContext, Tool, ToolCallEntry, ToolCallsMap, ToolInterface, ToolResult, WorkflowInterface, WorkflowMetadataInterface, } from '@loopstack/common'; import { AiGenerateToolBaseSchema } from '../schemas/ai-generate-tool-base.schema'; import { AiMessagesHelperService } from '../services'; import { AiProviderModelHelperService } from '../services'; import { AiToolsHelperService } from '../services'; export const AiGenerateTextSchema = AiGenerateToolBaseSchema.extend({ tools: z.array(z.string()).optional(), }).strict(); type AiGenerateTextArgsType = z.infer; @Tool({ config: { description: 'Generates text using a LLM', }, }) export class AiGenerateText implements ToolInterface { @Inject() private readonly aiMessagesHelperService: AiMessagesHelperService; @Inject() private readonly aiToolsHelperService: AiToolsHelperService; @Inject() private readonly aiProviderModelHelperService: AiProviderModelHelperService; @Input({ schema: AiGenerateTextSchema, }) args: AiGenerateTextArgsType; async execute( args: AiGenerateTextArgsType, ctx: RunContext, parent: WorkflowInterface, runtime: WorkflowMetadataInterface, ): Promise { const model = this.aiProviderModelHelperService.getProviderModel(args.llm); const options: { prompt?: string; messages: ModelMessage[]; tools?: Record; } = { messages: [], }; options.tools = args.tools ? this.aiToolsHelperService.getTools(args.tools, parent) : undefined; if (args.prompt) { options.messages.push({ role: 'user', content: args.prompt, } as ModelMessage); } else { const messages = this.aiMessagesHelperService.getMessages(runtime.documents, { messages: args.messages as unknown as UIMessage[], messagesSearchTag: args.messagesSearchTag, }); options.messages = await convertToModelMessages(messages, { tools: options.tools as ToolSet, }); } const { uiMessage, usage } = await this.handleGenerateText(model, options); const toolCalls = this.extractToolCalls(uiMessage); return { data: { ...uiMessage, ...(toolCalls ? { toolCalls } : {}), }, metadata: { usage }, }; } private extractToolCalls(message: UIMessage): ToolCallsMap | null { const toolCalls: ToolCallsMap = {}; for (const part of message.parts) { if (!('type' in part) || typeof part.type !== 'string' || !part.type.startsWith('tool-')) { continue; } const toolName = part.type.replace(/^tool-/, ''); toolCalls[toolName] = { id: (part as { toolCallId?: string }).toolCallId ?? '', name: toolName, input: (part as { input?: unknown }).input, } satisfies ToolCallEntry; } return Object.keys(toolCalls).length > 0 ? toolCalls : null; } private async handleGenerateText( model: LanguageModel, options: { prompt?: string; messages?: ModelMessage[]; tools?: Record; }, ): Promise<{ uiMessage: UIMessage; usage: { inputTokens: number; outputTokens: number } }> { const startTime = performance.now(); try { const result = streamText({ model, ...options, } as Parameters[0]); const uiMessage = await new Promise((resolve, reject) => { const stream = createUIMessageStream({ execute({ writer }) { writer.merge( result.toUIMessageStream({ sendReasoning: true, }), ); }, onFinish: (data) => { resolve(data.responseMessage); }, }); // Consume the stream to trigger execution void (async () => { try { const reader = stream.getReader(); while (true) { const { done } = await reader.read(); if (done) break; } } catch (error) { reject(error instanceof Error ? error : new Error(String(error))); } })(); }); const resultUsage = await result.usage; return { uiMessage, usage: { inputTokens: resultUsage.inputTokens ?? 0, outputTokens: resultUsage.outputTokens ?? 0, }, }; } catch (error) { const errorResponseTime = performance.now() - startTime; console.error(`Request failed after ${errorResponseTime}ms:`, error); throw error; } } }