diff --git a/src/server/sessions/piSessionService.test.ts b/src/server/sessions/piSessionService.test.ts index a024650..335bfd8 100644 --- a/src/server/sessions/piSessionService.test.ts +++ b/src/server/sessions/piSessionService.test.ts @@ -1,6 +1,8 @@ import { mkdtemp, rm, writeFile } from "node:fs/promises"; import { tmpdir } from "node:os"; import { join } from "node:path"; +import { createAssistantMessageEventStream, type AssistantMessage } from "@earendil-works/pi-ai"; +import type { StreamFn } from "@earendil-works/pi-agent-core"; import { AuthStorage, ModelRegistry } from "@earendil-works/pi-coding-agent"; import { describe, expect, it, vi } from "vitest"; import type { GlobalSessionEvent, SessionUiEvent } from "../../shared/apiTypes.js"; @@ -938,6 +940,42 @@ describe("PiSessionService", () => { await service.dispose(); }); + it("generates a session name for the first prompt via the session's agent.streamFn", async () => { + const model = testModel(); + const streamCalls: unknown[] = []; + const streamFn: StreamFn = (streamModel, context, options) => { + streamCalls.push({ streamModel, context, options }); + const stream = createAssistantMessageEventStream(); + const message: AssistantMessage = { + role: "assistant", + content: [{ type: "text", text: "Fix login bug" }], + api: "anthropic-messages", + provider: "anthropic", + model: model.id, + usage: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, totalTokens: 0, cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 } }, + stopReason: "stop", + timestamp: Date.now(), + }; + stream.push({ type: "done", reason: "stop", message }); + stream.end(message); + return stream; + }; + const hub = new CapturingSessionEventHub(); + const fake = fakeRuntime("name-session", { model, agent: { streamFn } }); + const service = new PiSessionService(hub, { + createAgentRuntime: runtimeCreator(fake.runtime), + sessionManager: sessionGateway([sessionRecord("name-session")]), + heartbeatIntervalMs: 60_000, + }); + + await service.prompt(sessionRef("name-session"), "Please fix the login bug"); + await vi.waitFor(() => { expect(fake.session.sessionName).toBe("Fix login bug"); }); + + expect(streamCalls).toHaveLength(1); + expect(hub.sessionEvents.some(({ event }) => event.type === "session.name" && event.name === "Fix login bug")).toBe(true); + await service.dispose(); + }); + it("includes queued message details in session status", async () => { const fake = fakeRuntime("status-session", { messages: [{ role: "user", content: "hello" }, { role: "assistant", content: "hi" }], diff --git a/src/server/sessions/piSessionService.ts b/src/server/sessions/piSessionService.ts index e38ffb2..37e5f03 100644 --- a/src/server/sessions/piSessionService.ts +++ b/src/server/sessions/piSessionService.ts @@ -1707,7 +1707,7 @@ export class PiSessionService { const model = session.model; if (model === undefined) return; - void generateShortSessionName(this.modelRegistry, model, firstMessage).then((name) => { + void generateShortSessionName(session.agent.streamFn, model, firstMessage).then((name) => { this.applyGeneratedSessionName(session, name ?? fallbackSessionName(firstMessage)); }).catch(() => { this.applyGeneratedSessionName(session, fallbackSessionName(firstMessage)); diff --git a/src/server/sessions/sessionNameGenerator.test.ts b/src/server/sessions/sessionNameGenerator.test.ts index 5a93d50..12078cb 100644 --- a/src/server/sessions/sessionNameGenerator.test.ts +++ b/src/server/sessions/sessionNameGenerator.test.ts @@ -1,7 +1,70 @@ +import type { Api, AssistantMessage, Model } from "@earendil-works/pi-ai"; +import { createAssistantMessageEventStream } from "@earendil-works/pi-ai"; +import type { StreamFn } from "@earendil-works/pi-agent-core"; import { describe, expect, it } from "vitest"; -import { cleanSessionName, fallbackSessionName } from "./sessionNameGenerator.js"; +import { cleanSessionName, fallbackSessionName, generateShortSessionName } from "./sessionNameGenerator.js"; + +function fakeModel(): Model { + return { id: "fake-model", name: "Fake Model", api: "anthropic-messages", provider: "anthropic", baseUrl: "https://example.test", reasoning: false, input: ["text"], cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0 }, contextWindow: 1000, maxTokens: 100 }; +} + +function fakeAssistantMessage(overrides: Partial = {}): AssistantMessage { + return { + role: "assistant", + content: [], + api: "anthropic-messages", + provider: "anthropic", + model: "fake-model", + usage: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, totalTokens: 0, cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 } }, + stopReason: "stop", + timestamp: Date.now(), + ...overrides, + }; +} + +function streamThatCompletes(text: string): StreamFn { + return () => { + const stream = createAssistantMessageEventStream(); + const message = fakeAssistantMessage({ content: [{ type: "text", text }] }); + stream.push({ type: "done", reason: "stop", message }); + stream.end(message); + return stream; + }; +} + +function streamThatErrors(): StreamFn { + return () => { + const stream = createAssistantMessageEventStream(); + const message = fakeAssistantMessage({ stopReason: "error", errorMessage: "boom" }); + stream.push({ type: "error", reason: "error", error: message }); + stream.end(message); + return stream; + }; +} describe("sessionNameGenerator", () => { + it("generates a session name by calling the injected streamFn", async () => { + const calls: unknown[] = []; + const stream = streamThatCompletes('Title: "Fix the bug"'); + const streamFn: StreamFn = (model, context, options) => { + calls.push({ model, context, options }); + return stream(model, context, options); + }; + + const name = await generateShortSessionName(streamFn, fakeModel(), "Please fix the login bug"); + + expect(name).toBe("Fix the bug"); + expect(calls).toHaveLength(1); + }); + + it("returns undefined when the stream reports an error", async () => { + const streamFn = streamThatErrors(); + + const name = await generateShortSessionName(streamFn, fakeModel(), "Please fix the login bug"); + + expect(name).toBeUndefined(); + }); + it("cleans model-generated titles", () => { expect(cleanSessionName('Title: "Fix Session Naming."\nextra')).toBe("Fix Session Naming"); }); diff --git a/src/server/sessions/sessionNameGenerator.ts b/src/server/sessions/sessionNameGenerator.ts index 77c7c15..b95a5d7 100644 --- a/src/server/sessions/sessionNameGenerator.ts +++ b/src/server/sessions/sessionNameGenerator.ts @@ -1,33 +1,13 @@ -import type { Api, AssistantMessage, AssistantMessageEventStream, Context, Model, SimpleStreamOptions } from "@earendil-works/pi-ai"; -import type { ModelRegistry } from "@earendil-works/pi-coding-agent"; +import type { Api, AssistantMessage, Model } from "@earendil-works/pi-ai"; +import type { StreamFn } from "@earendil-works/pi-agent-core"; const SESSION_NAME_TIMEOUT_MS = 10_000; const SESSION_NAME_MAX_INPUT_CHARS = 4_000; const SESSION_NAME_MAX_LENGTH = 60; const FALLBACK_SESSION_NAME_MAX_WORDS = 6; -const PI_AI_COMPAT_MODULE = ["@earendil-works/pi-ai", "compat"].join("/"); -interface SessionNameApiProvider { - streamSimple(model: Model, context: Context, options?: SimpleStreamOptions): AssistantMessageEventStream; -} - -interface PiAiProviderRegistryModule { - getApiProvider?: (api: Api) => SessionNameApiProvider | undefined; -} - -type ModuleImporter = (specifier: string) => Promise; - -let piAiProviderRegistryModulePromise: Promise | undefined; - -export async function generateShortSessionName(modelRegistry: ModelRegistry, model: Model, firstMessage: string): Promise { - const providerRegistry = await getPiAiProviderRegistryModule(); - const provider = providerRegistry.getApiProvider?.(model.api); - if (provider === undefined) return undefined; - - const auth = await modelRegistry.getApiKeyAndHeaders(model); - if (!auth.ok) return undefined; - - const stream = provider.streamSimple( +export async function generateShortSessionName(streamFn: StreamFn, model: Model, firstMessage: string): Promise { + const stream = await streamFn( model, { systemPrompt: "Generate a concise title for a coding-agent chat session. Return only the title, with no quotes or punctuation wrapper.", @@ -41,8 +21,6 @@ export async function generateShortSessionName(modelRegistry: maxTokens: 24, reasoning: "minimal", signal: AbortSignal.timeout(SESSION_NAME_TIMEOUT_MS), - ...(auth.apiKey === undefined ? {} : { apiKey: auth.apiKey }), - ...(auth.headers === undefined ? {} : { headers: auth.headers }), }, ); @@ -81,42 +59,6 @@ export function cleanSessionName(value: string): string | undefined { return title === "" ? undefined : title; } -async function getPiAiProviderRegistryModule(importer: ModuleImporter = (specifier) => import(specifier)): Promise { - piAiProviderRegistryModulePromise ??= loadPiAiProviderRegistryModule(importer); - return piAiProviderRegistryModulePromise; -} - -async function loadPiAiProviderRegistryModule(importer: ModuleImporter): Promise { - const compatModule = await importOptionalPiAiModule(PI_AI_COMPAT_MODULE, importer); - if (hasGetApiProvider(compatModule)) return compatModule; - - const rootModule = await importer("@earendil-works/pi-ai"); - if (hasGetApiProvider(rootModule)) return rootModule; - return {}; -} - -async function importOptionalPiAiModule(specifier: string, importer: ModuleImporter): Promise { - try { - return await importer(specifier); - } catch (error) { - if (isModuleUnavailableError(error)) return undefined; - throw error; - } -} - -function hasGetApiProvider(moduleValue: unknown): moduleValue is PiAiProviderRegistryModule { - return typeof moduleValue === "object" - && moduleValue !== null - && "getApiProvider" in moduleValue - && typeof moduleValue.getApiProvider === "function"; -} - -function isModuleUnavailableError(error: unknown): boolean { - if (!(error instanceof Error)) return false; - const code = "code" in error ? error.code : undefined; - return code === "ERR_MODULE_NOT_FOUND" || code === "ERR_PACKAGE_PATH_NOT_EXPORTED"; -} - function textFromAssistant(message: AssistantMessage): string { return message.content .filter((part) => part.type === "text")