refactor: enhance session management and LLM processing

- Introduce a new `currentStreamPromise` to track ongoing LLM processing, allowing for better handling of stream interruptions and cleanup.
- Update message handling to utilize `ModelMessage` for improved type safety and structure.
- Refactor `processWithLLM` to integrate new streaming capabilities with tool execution and error handling.
- Enhance logging for better traceability during audio playback and message processing.
- Remove deprecated `streamLLM` function and replace it with a more modular approach using `getAllTools`.

These changes improve the robustness of the voice assistant's session management and streamline the integration of LLM processing.
This commit is contained in:
Mohamed Boudra
2025-10-20 16:14:25 +02:00
parent af67235997
commit c5b24d3dfc
3 changed files with 229 additions and 227 deletions

View File

@@ -1,8 +1,5 @@
import { stepCountIs, streamText, tool } from "ai";
import { createOpenRouter } from "@openrouter/ai-sdk-provider";
import { tool, experimental_createMCPClient } from "ai";
import { z } from "zod";
import invariant from "tiny-invariant";
import { experimental_createMCPClient } from "ai";
import { InMemoryTransport } from "@modelcontextprotocol/sdk/inMemory.js";
import { StdioClientTransport } from "@modelcontextprotocol/sdk/client/stdio.js";
import { createTerminalMcpServer } from "../terminal-mcp/index.js";
@@ -154,125 +151,14 @@ const manualTools = {
};
/**
* Message interface for conversation
* Get all tools (MCP + manual) for LLM
*/
export interface Message {
role: "system" | "user" | "assistant";
content: string;
}
/**
* Streaming LLM parameters
*/
export interface StreamLLMParams {
systemPrompt: string;
messages: Message[];
abortSignal?: AbortSignal;
onChunk?: (chunk: string) => void | Promise<void>;
onTextSegment?: (segment: string) => void;
onToolCall?: (
toolCallId: string,
toolName: string,
args: any
) => Promise<void>;
onToolResult?: (toolCallId: string, toolName: string, result: any) => void;
onToolError?: (
toolCallId: string,
toolName: string,
error: unknown
) => void | Promise<void>;
onError?: (error: unknown) => void | Promise<void>;
onFinish?: (fullText: string) => void | Promise<void>;
}
/**
* Stream LLM response with automatic tool execution
*/
export async function streamLLM(params: StreamLLMParams): Promise<string> {
invariant(process.env.OPENROUTER_API_KEY, "OPENROUTER_API_KEY is required");
const openrouter = createOpenRouter({
apiKey: process.env.OPENROUTER_API_KEY,
});
// Get all MCP tools and merge with manual tools
export async function getAllTools(): Promise<Record<string, any>> {
const mcpTools = await getMcpTools();
const allTools = {
return {
...mcpTools,
...manualTools,
};
const result = await streamText({
model: openrouter("anthropic/claude-haiku-4.5"),
system: params.systemPrompt,
messages: params.messages,
tools: allTools,
abortSignal: params.abortSignal,
onChunk: async ({ chunk }) => {
// console.log("onChunk", chunk);
if (chunk.type === "text-delta") {
// Accumulate text in buffer
textBuffer += chunk.text;
fullText += chunk.text;
params.onChunk?.(chunk.text);
} else if (chunk.type === "tool-call") {
// Flush accumulated text as a segment before tool call
flushTextBuffer();
// Emit tool call event
if (params.onToolCall) {
await params.onToolCall(
chunk.toolCallId,
chunk.toolName,
chunk.input
);
}
} else if (chunk.type === "tool-result") {
// Emit tool result event
if (params.onToolResult) {
params.onToolResult(chunk.toolCallId, chunk.toolName, chunk.output);
}
}
},
onError: async (error) => {
// Emit general stream error
if (params.onError) {
await params.onError(error);
}
},
stopWhen: stepCountIs(10),
});
let fullText = "";
let textBuffer = "";
function flushTextBuffer() {
if (textBuffer.length > 0 && params.onTextSegment) {
params.onTextSegment(textBuffer);
}
textBuffer = "";
}
for await (const part of result.fullStream) {
// console.log("part", part);
// Handle tool-error chunks (not available in onChunk callback)
if (part.type === "tool-error") {
if (params.onToolError) {
await params.onToolError(part.toolCallId, part.toolName, part.error);
}
}
}
// Flush any remaining text at the end
flushTextBuffer();
if (params.onFinish) {
await params.onFinish(fullText);
}
return fullText;
}
/**

View File

@@ -50,14 +50,18 @@ export class TTSManager {
});
console.log(
`[TTS-Manager ${this.sessionId}] ${new Date().toISOString()} Sent audio ${audioId}, waiting for playback...`
`[TTS-Manager ${
this.sessionId
}] ${new Date().toISOString()} Sent audio ${audioId}, waiting for playback...`
);
// Wait for playback confirmation
await playbackPromise;
console.log(
`[TTS-Manager ${this.sessionId}] ${new Date().toISOString()} Audio ${audioId} playback confirmed`
`[TTS-Manager ${
this.sessionId
}] ${new Date().toISOString()} Audio ${audioId} playback confirmed`
);
}

View File

@@ -2,12 +2,16 @@ import { v4 as uuidv4 } from "uuid";
import { readFile } from "fs/promises";
import { exec } from "child_process";
import { promisify } from "util";
import invariant from "tiny-invariant";
import { streamText, stepCountIs } from "ai";
import type { ModelMessage } from "@ai-sdk/provider-utils";
import { createOpenRouter } from "@openrouter/ai-sdk-provider";
import type {
SessionInboundMessage,
SessionOutboundMessage,
} from "./messages.js";
import { getSystemPrompt } from "./agent/system-prompt.js";
import { streamLLM, type Message } from "./agent/llm-openai.js";
import { getAllTools } from "./agent/llm-openai.js";
import { TTSManager } from "./agent/tts-manager.js";
import { STTManager } from "./agent/stt-manager.js";
@@ -40,6 +44,7 @@ export class Session {
// State machine
private abortController: AbortController;
private processingPhase: ProcessingPhase = "idle";
private currentStreamPromise: Promise<void> | null = null;
// Audio buffering for interruption handling
private pendingAudioSegments: Array<{ audio: Buffer; format: string }> = [];
@@ -47,7 +52,7 @@ export class Session {
private audioBuffer: { chunks: Buffer[]; format: string } | null = null;
// Conversation history
private messages: Message[] = [];
private messages: ModelMessage[] = [];
// Per-session managers
private readonly ttsManager: TTSManager;
@@ -101,7 +106,10 @@ export class Session {
break;
}
} catch (error: any) {
console.error(`[Session ${this.clientId}] Error handling message:`, error);
console.error(
`[Session ${this.clientId}] Error handling message:`,
error
);
this.emit({
type: "activity_log",
payload: {
@@ -118,9 +126,17 @@ export class Session {
* Handle text message from user
*/
private async handleUserText(text: string): Promise<void> {
// Create new abort controller
// Abort any in-progress stream immediately
this.createAbortController();
// Wait for aborted stream to finish cleanup (save partial response)
if (this.currentStreamPromise) {
console.log(
`[Session ${this.clientId}] Waiting for aborted stream to finish cleanup`
);
await this.currentStreamPromise;
}
// Emit user message activity log
this.emit({
type: "activity_log",
@@ -136,7 +152,8 @@ export class Session {
this.messages.push({ role: "user", content: text });
// Process through LLM (TTS enabled for text input)
await this.processWithLLM(true);
this.currentStreamPromise = this.processWithLLM(true);
await this.currentStreamPromise;
}
/**
@@ -218,7 +235,17 @@ export class Session {
* Process audio through STT and then LLM
*/
private async processAudio(audio: Buffer, format: string): Promise<void> {
// Abort any in-progress stream immediately
this.createAbortController();
// Wait for aborted stream to finish cleanup (save partial response)
if (this.currentStreamPromise) {
console.log(
`[Session ${this.clientId}] Waiting for aborted stream to finish cleanup`
);
await this.currentStreamPromise;
}
this.setPhase("transcribing");
this.emit({
@@ -273,7 +300,8 @@ export class Session {
// Set phase to LLM and process
this.setPhase("llm");
await this.processWithLLM(true); // Enable TTS for voice input
this.currentStreamPromise = this.processWithLLM(true); // Enable TTS for voice input
await this.currentStreamPromise;
this.setPhase("idle");
} catch (error: any) {
this.setPhase("idle");
@@ -296,90 +324,129 @@ export class Session {
private async processWithLLM(enableTTS: boolean): Promise<void> {
let assistantResponse = "";
let pendingTTS: Promise<void> | null = null;
let textBuffer = "";
const flushTextBuffer = () => {
if (textBuffer.length > 0) {
// TTS handling
if (enableTTS) {
pendingTTS = this.ttsManager.generateAndWaitForPlayback(
textBuffer,
(msg) => this.emit(msg)
);
}
// Emit activity log
this.emit({
type: "activity_log",
payload: {
id: uuidv4(),
timestamp: new Date(),
type: "assistant",
content: textBuffer,
},
});
}
textBuffer = "";
};
try {
assistantResponse = await streamLLM({
systemPrompt: getSystemPrompt(),
invariant(
process.env.OPENROUTER_API_KEY,
"OPENROUTER_API_KEY is required"
);
const openrouter = createOpenRouter({
apiKey: process.env.OPENROUTER_API_KEY,
});
const allTools = await getAllTools();
const result = await streamText({
model: openrouter("anthropic/claude-haiku-4.5"),
system: getSystemPrompt(),
messages: this.messages,
tools: allTools,
abortSignal: this.abortController.signal,
onTextSegment: (segment) => {
if (enableTTS) {
// Create TTS promise (don't await yet)
pendingTTS = this.ttsManager.generateAndWaitForPlayback(
segment,
(msg) => this.emit(msg)
);
onFinish: async (event) => {
for (const step of event.steps) {
this.messages.push(...step.response.messages);
}
// Emit activity log
this.emit({
type: "activity_log",
payload: {
id: uuidv4(),
timestamp: new Date(),
type: "assistant",
content: segment,
},
});
console.log(
`[Session ${this.clientId}] onFinish - saved message with ${event.steps.length} steps`
);
},
onChunk: async (chunk) => {
this.emit({
type: "assistant_chunk",
payload: { chunk },
});
},
onToolCall: async (toolCallId, toolName, args) => {
// Wait for pending TTS before executing tool
if (pendingTTS) {
console.log(
`[Session ${this.clientId}] Waiting for TTS before executing ${toolName}`
);
await pendingTTS;
pendingTTS = null;
onAbort: async (event) => {
for (const step of event.steps) {
this.messages.push(...step.response.messages);
}
console.log(
`[Session ${this.clientId}] onAbort - saved message with ${event.steps.length} steps`
);
},
onChunk: async ({ chunk }) => {
if (chunk.type === "text-delta") {
// Accumulate text in buffer
textBuffer += chunk.text;
assistantResponse += chunk.text;
// Handle present_artifact tool specially
if (toolName === "present_artifact") {
await this.handlePresentArtifact(toolCallId, args as PresentArtifactArgs);
// Emit chunk for UI streaming
this.emit({
type: "assistant_chunk",
payload: { chunk: chunk.text },
});
} else if (chunk.type === "tool-call") {
// Flush accumulated text as a segment before tool call
flushTextBuffer();
// Wait for pending TTS before executing tool
if (pendingTTS) {
console.log(
`[Session ${this.clientId}] Waiting for TTS before executing ${chunk.toolName}`
);
await pendingTTS;
}
// Handle present_artifact tool specially
if (chunk.toolName === "present_artifact") {
await this.handlePresentArtifact(
chunk.toolCallId,
chunk.input as PresentArtifactArgs
);
}
// Emit tool call activity log
this.emit({
type: "activity_log",
payload: {
id: chunk.toolCallId,
timestamp: new Date(),
type: "tool_call",
content: `Calling ${chunk.toolName}`,
metadata: {
toolCallId: chunk.toolCallId,
toolName: chunk.toolName,
arguments: chunk.input,
},
},
});
} else if (chunk.type === "tool-result") {
// Emit tool result event
this.emit({
type: "activity_log",
payload: {
id: chunk.toolCallId,
timestamp: new Date(),
type: "tool_result",
content: `Tool ${chunk.toolName} completed`,
metadata: {
toolCallId: chunk.toolCallId,
toolName: chunk.toolName,
result: chunk.output,
},
},
});
}
// Emit tool call activity log
this.emit({
type: "activity_log",
payload: {
id: toolCallId,
timestamp: new Date(),
type: "tool_call",
content: `Calling ${toolName}`,
metadata: { toolCallId, toolName, arguments: args },
},
});
},
onToolResult: (toolCallId, toolName, result) => {
this.emit({
type: "activity_log",
payload: {
id: toolCallId,
timestamp: new Date(),
type: "tool_result",
content: `Tool ${toolName} completed`,
metadata: { toolCallId, toolName, result },
},
});
},
onToolError: async (toolCallId, toolName, error) => {
this.emit({
type: "activity_log",
payload: {
id: toolCallId,
timestamp: new Date(),
type: "error",
content: `Tool ${toolName} failed: ${
error instanceof Error ? error.message : String(error)
}`,
metadata: { toolCallId, toolName, error },
},
});
},
onError: async (error) => {
this.emit({
@@ -394,16 +461,37 @@ export class Session {
},
});
},
onFinish: async () => {
// Don't wait for TTS here
},
stopWhen: stepCountIs(10),
});
// Add assistant response to conversation IMMEDIATELY after stream completes
this.messages.push({
role: "assistant",
content: assistantResponse,
});
// Consume the fullStream to handle tool-error chunks
for await (const part of result.fullStream) {
if (part.type === "tool-error") {
this.emit({
type: "activity_log",
payload: {
id: part.toolCallId,
timestamp: new Date(),
type: "error",
content: `Tool ${part.toolName} failed: ${
part.error instanceof Error
? part.error.message
: String(part.error)
}`,
metadata: {
toolCallId: part.toolCallId,
toolName: part.toolName,
error: part.error,
},
},
});
}
}
// Flush any remaining text at the end
flushTextBuffer();
// Note: Message is saved by onFinish callback with proper tool calls
// Now wait for any pending TTS, but don't fail if it times out
if (pendingTTS) {
@@ -417,26 +505,40 @@ export class Session {
}
}
} catch (error) {
// If stream failed or was aborted, still save any partial response
if (assistantResponse) {
this.messages.push({
role: "assistant",
content: assistantResponse,
// Note: Partial messages are saved by onAbort callback with proper tool calls
// Check if this is an abort error
const isAbortError =
error instanceof Error && error.name === "AbortError";
// Only emit error log for non-abort errors
if (!isAbortError) {
this.emit({
type: "activity_log",
payload: {
id: uuidv4(),
timestamp: new Date(),
type: "error",
content: `Error: ${
error instanceof Error ? error.message : String(error)
}`,
},
});
}
this.emit({
type: "activity_log",
payload: {
id: uuidv4(),
timestamp: new Date(),
type: "error",
content: `Error: ${
error instanceof Error ? error.message : String(error)
}`,
},
});
// Don't re-throw abort errors (they're expected during interruptions)
if (isAbortError) {
console.log(
`[Session ${this.clientId}] Stream aborted (partial response saved)`
);
return;
}
// Re-throw unexpected errors
throw error;
} finally {
// Clear the stream promise tracker
this.currentStreamPromise = null;
}
}
@@ -473,7 +575,9 @@ export class Session {
`[Session ${this.clientId}] Failed to resolve artifact source:`,
error
);
content = `[Error: ${error instanceof Error ? error.message : String(error)}]`;
content = `[Error: ${
error instanceof Error ? error.message : String(error)
}]`;
isBase64 = false;
}
@@ -511,10 +615,18 @@ export class Session {
);
if (this.processingPhase === "llm") {
// Already in LLM phase - abort immediately
// Already in LLM phase - abort and wait for cleanup
this.abortController.abort();
console.log(`[Session ${this.clientId}] Aborted LLM processing`);
// Wait for stream to finish saving partial response
if (this.currentStreamPromise) {
console.log(
`[Session ${this.clientId}] Waiting for stream cleanup after abort`
);
await this.currentStreamPromise;
}
// Reset phase to idle
this.setPhase("idle");