Refactor: Move config and types to separate files, enhance tool call handling.

This commit is contained in:
uxname committed 2025-05-24 15:34:35 +03:00
1 parent 6047368080
commit 65a8898d14
1 file changed
+81 -58
+81 -58
View File
@@ -1,4 +1,5 @@
import { import {
AIMessage,
HumanMessage, HumanMessage,
SystemMessage, SystemMessage,
ToolMessage, ToolMessage,
@@ -27,7 +28,6 @@ class MockMcpToolsService {
description: 'Возвращает введенную строку.', description: 'Возвращает введенную строку.',
schema: z.object({ input: z.string().describe('Строка для эхо.') }), schema: z.object({ input: z.string().describe('Строка для эхо.') }),
func: async (input: { input: string }) => { func: async (input: { input: string }) => {
console.log(`Вызов инструмента 'echo' с входом: ${input.input}`);
return input.input; return input.input;
}, },
}); });
@@ -37,7 +37,6 @@ class MockMcpToolsService {
description: 'Возвращает текущую дату в формате ISO.', description: 'Возвращает текущую дату в формате ISO.',
schema: z.object({}), schema: z.object({}),
func: async () => { func: async () => {
console.log('Вызов инструмента "get_current_date"');
return new Date().toISOString(); return new Date().toISOString();
}, },
}); });
@@ -48,13 +47,21 @@ class MockMcpToolsService {
class PipelineExecutor { class PipelineExecutor {
private readonly mcpToolsService: MockMcpToolsService; private readonly mcpToolsService: MockMcpToolsService;
private cachedTools: DynamicStructuredTool[] | null = null;
constructor() { constructor() {
this.mcpToolsService = new MockMcpToolsService(); this.mcpToolsService = new MockMcpToolsService();
} }
private async _getOrLoadTools(): Promise<DynamicStructuredTool[]> {
if (!this.cachedTools) {
this.cachedTools = await this.mcpToolsService.getTools();
}
return this.cachedTools;
}
private async initChain(openAIApiKey: string): Promise<Runnable> { private async initChain(openAIApiKey: string): Promise<Runnable> {
const tools = await this.mcpToolsService.getTools(); const tools = await this._getOrLoadTools();
const llm = new ChatOpenAI({ const llm = new ChatOpenAI({
model: 'gpt-4o', model: 'gpt-4o',
temperature: 0.2, temperature: 0.2,
@@ -68,10 +75,8 @@ class PipelineExecutor {
variables: Record<string, string> = {}, variables: Record<string, string> = {},
): (HumanMessage | SystemMessage)[] { ): (HumanMessage | SystemMessage)[] {
return [ return [
...systemPrompts.map((prompt) => ...systemPrompts.map(
Object.keys(variables).length (prompt) => new SystemMessage(Handlebars.compile(prompt)(variables)),
? new SystemMessage(Handlebars.compile(prompt)(variables))
: new SystemMessage(prompt),
), ),
new SystemMessage(`Текущая дата и время: ${new Date().toISOString()}`), new SystemMessage(`Текущая дата и время: ${new Date().toISOString()}`),
]; ];
@@ -81,73 +86,91 @@ class PipelineExecutor {
prompt: string, prompt: string,
variables: Record<string, string> = {}, variables: Record<string, string> = {},
): string { ): string {
return Object.keys(variables).length return Handlebars.compile(prompt)(variables);
? Handlebars.compile(prompt)(variables)
: prompt;
} }
async executeChain(params: { async executeChain(options: {
pipeline: Pipeline; pipeline: Pipeline;
variables?: Record<string, string>; variables?: Record<string, string>;
systemPrompts?: string[]; systemPrompts?: string[];
openAIApiKey: string; openAIApiKey: string;
}): Promise<string[]> { }): Promise<string[]> {
const variables = params.variables ?? {}; const variables = options.variables ?? {};
const systemPrompts = params.systemPrompts ?? []; const systemPrompts = options.systemPrompts ?? [];
const chain = await this.initChain(params.openAIApiKey); const chain = await this.initChain(options.openAIApiKey);
const initialMessages = this.buildMessages( const initialMessages = this.buildMessages(
[params.pipeline.systemPrompt, ...systemPrompts], [options.pipeline.systemPrompt, ...systemPrompts],
variables, variables,
); );
const results: string[] = []; const results: string[] = [];
for (const { prompt } of params.pipeline.steps) { for (const { prompt } of options.pipeline.steps) {
const renderedPrompt = this.renderPrompt(prompt, variables); const renderedPrompt = this.renderPrompt(prompt, variables);
const currentMessages = [ const currentMessages: (
...initialMessages, | HumanMessage
new HumanMessage(renderedPrompt), | SystemMessage
]; | AIMessage
let response = await chain.invoke(currentMessages); | ToolMessage
)[] = [...initialMessages, new HumanMessage(renderedPrompt)];
while (response.tool_calls && response.tool_calls.length > 0) { const finalResponse = await this._resolveToolCalls(
const toolCalls = response.tool_calls; chain,
currentMessages.push(response); currentMessages,
);
for (const toolCall of toolCalls) { results.push(finalResponse?.content.toString() ?? '<empty>');
try {
const tools = await this.mcpToolsService.getTools();
const tool = tools.find((t) => t.name === toolCall.name);
if (!tool) {
throw new Error(`Tool ${toolCall.name} not found`);
}
const toolResult = await tool.func(toolCall.args);
currentMessages.push(
new ToolMessage({
tool_call_id: toolCall.id,
content: JSON.stringify(toolResult),
}),
);
} catch (error) {
console.error(
`Ошибка выполнения инструмента ${toolCall.name}:`,
error,
);
currentMessages.push(
new ToolMessage({
tool_call_id: toolCall.id,
content: `Error: ${error.message}`,
}),
);
}
}
response = await chain.invoke(currentMessages);
}
results.push(response?.content ?? '<empty>');
} }
return results; return results;
} }
private async _resolveToolCalls(
chain: Runnable,
messages: (HumanMessage | SystemMessage | AIMessage | ToolMessage)[],
): Promise<AIMessage> {
let response: AIMessage = await chain.invoke(messages);
while (response.tool_calls && response.tool_calls.length > 0) {
messages.push(response);
for (const toolCall of response.tool_calls) {
try {
const tools = await this._getOrLoadTools();
const tool = tools.find((t) => t.name === toolCall.name);
if (!tool) {
console.error(`Инструмент ${toolCall.name} не найден.`);
messages.push(
new ToolMessage({
tool_call_id: toolCall.id!,
content: `Ошибка: Инструмент ${toolCall.name} не найден.`,
}),
);
continue;
}
const toolResult = await tool.func(toolCall.args);
messages.push(
new ToolMessage({
tool_call_id: toolCall.id!,
content: JSON.stringify(toolResult),
}),
);
} catch (error) {
console.error(
`Ошибка выполнения инструмента ${toolCall.name}:`,
error,
);
messages.push(
new ToolMessage({
tool_call_id: toolCall.id!,
content: `Ошибка: ${error.message}`,
}),
);
}
}
response = await chain.invoke(messages);
}
return response;
}
} }
async function main() { async function main() {
@@ -181,7 +204,7 @@ async function main() {
name: 'Получение даты', name: 'Получение даты',
prompt: 'Какая сейчас дата? Используй инструмент `get_current_date`.', prompt: 'Какая сейчас дата? Используй инструмент `get_current_date`.',
}, },
{ name: 'Завершение', prompt: 'Напиши резюме, что ты сделал' }, { name: 'Завершение', prompt: 'Напиши резюме нашего с тобой диалога' },
], ],
}; };