From 82ba2e0490b709e65c1c63a86a416b8c940a5d8a Mon Sep 17 00:00:00 2001 From: Federico Jaramillo Martinez Date: Fri, 5 Jun 2026 10:59:01 +0200 Subject: [PATCH] fix: validate prompt API payloads Fixes #11 --- .changeset/guard-missing-prompt-text.md | 5 ++ src/server/sessions/piSessionService.test.ts | 14 +++++ src/server/sessions/piSessionService.ts | 25 ++++++-- .../sessions/sessionNameGenerator.test.ts | 4 ++ src/server/sessions/sessionNameGenerator.ts | 4 +- src/server/sessions/sessionRoutes.test.ts | 58 +++++++++++++++++++ src/server/sessions/sessionRoutes.ts | 9 ++- 7 files changed, 110 insertions(+), 9 deletions(-) create mode 100644 .changeset/guard-missing-prompt-text.md create mode 100644 src/server/sessions/sessionRoutes.test.ts diff --git a/.changeset/guard-missing-prompt-text.md b/.changeset/guard-missing-prompt-text.md new file mode 100644 index 0000000..c094f79 --- /dev/null +++ b/.changeset/guard-missing-prompt-text.md @@ -0,0 +1,5 @@ +--- +"@jmfederico/pi-web": patch +--- + +Prevent malformed session prompt API calls from crashing the session daemon. diff --git a/src/server/sessions/piSessionService.test.ts b/src/server/sessions/piSessionService.test.ts index eb23fc1..f04535e 100644 --- a/src/server/sessions/piSessionService.test.ts +++ b/src/server/sessions/piSessionService.test.ts @@ -366,6 +366,20 @@ describe("PiSessionService", () => { await service.dispose(); }); + it("rejects malformed prompt text before opening the runtime", async () => { + const fake = fakeRuntime("prompt-session"); + const service = new PiSessionService(new CapturingSessionEventHub(), { + createAgentRuntime: runtimeCreator(fake.runtime), + sessionManager: sessionGateway([sessionRecord("prompt-session")]), + heartbeatIntervalMs: 60_000, + }); + + await expect(service.prompt("prompt-session", undefined)).rejects.toThrow("Prompt text is required"); + + expect(fake.calls.prompt).toEqual([]); + 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 cdb29ed..26b6df8 100644 --- a/src/server/sessions/piSessionService.ts +++ b/src/server/sessions/piSessionService.ts @@ -41,6 +41,17 @@ interface QueuedPrompt { text: string; } +function requirePromptText(value: unknown): string { + if (typeof value !== "string") throw new Error("Prompt text is required"); + return value; +} + +function parsePromptStreamingBehavior(value: unknown): QueuedPromptKind | undefined { + if (value === undefined) return undefined; + if (value === "steer" || value === "followUp") return value; + throw new Error('Prompt streamingBehavior must be "steer" or "followUp"'); +} + type SessionArchiveRepository = Pick; interface PiSessionListEntry { id: string; @@ -360,22 +371,24 @@ export class PiSessionService { return commands.sort((a, b) => a.name.localeCompare(b.name)); } - async prompt(sessionId: string, text: string, streamingBehavior?: "steer" | "followUp"): Promise { + async prompt(sessionId: string, text: unknown, streamingBehavior?: unknown): Promise { + const promptText = requirePromptText(text); + const requestedBehavior = parsePromptStreamingBehavior(streamingBehavior); await this.assertWritable(sessionId); const session = await this.getOrOpen(sessionId); - this.maybeGenerateSessionName(session, text); + this.maybeGenerateSessionName(session, promptText); const isQueued = session.isStreaming || session.isCompacting; - const behavior = isQueued ? streamingBehavior ?? "followUp" : undefined; - if (isQueued && this.hasQueuedMessageText(session, text)) { + const behavior = isQueued ? requestedBehavior ?? "followUp" : undefined; + if (isQueued && this.hasQueuedMessageText(session, promptText)) { this.publishActivity(session, "duplicate queued message ignored", "active"); this.publishStatus(session); return; } if (session.isCompacting) { - this.enqueuePromptDuringCompaction(session, text, behavior ?? "followUp"); + this.enqueuePromptDuringCompaction(session, promptText, behavior ?? "followUp"); return; } - void this.submitPrompt(session, text, behavior); + void this.submitPrompt(session, promptText, behavior); } private submitPrompt(session: PiAgentSession, text: string, behavior: QueuedPromptKind | undefined): Promise { diff --git a/src/server/sessions/sessionNameGenerator.test.ts b/src/server/sessions/sessionNameGenerator.test.ts index be26d68..5a93d50 100644 --- a/src/server/sessions/sessionNameGenerator.test.ts +++ b/src/server/sessions/sessionNameGenerator.test.ts @@ -15,4 +15,8 @@ describe("sessionNameGenerator", () => { expect(fallbackSessionName('\nDo x\n\n\nCheck the UI now')) .toBe("Check the UI now"); }); + + it("skips fallback names when the first request is missing", () => { + expect(fallbackSessionName(undefined)).toBeUndefined(); + }); }); diff --git a/src/server/sessions/sessionNameGenerator.ts b/src/server/sessions/sessionNameGenerator.ts index 0c5ed07..d214ea3 100644 --- a/src/server/sessions/sessionNameGenerator.ts +++ b/src/server/sessions/sessionNameGenerator.ts @@ -43,7 +43,9 @@ export async function generateShortSessionName(modelRegistry: return cleanSessionName(finalMessage === undefined ? streamedText : textFromAssistant(finalMessage)); } -export function fallbackSessionName(firstMessage: string): string | undefined { +export function fallbackSessionName(firstMessage: unknown): string | undefined { + if (typeof firstMessage !== "string") return undefined; + return cleanSessionName(firstMessage .replace(/[\s\S]*?<\/skill>/g, "") .replace(/```[\s\S]*?```/g, " ") diff --git a/src/server/sessions/sessionRoutes.test.ts b/src/server/sessions/sessionRoutes.test.ts new file mode 100644 index 0000000..4e9c438 --- /dev/null +++ b/src/server/sessions/sessionRoutes.test.ts @@ -0,0 +1,58 @@ +import Fastify, { type FastifyInstance } from "fastify"; +import fastifyWebsocket from "@fastify/websocket"; +import { afterEach, beforeEach, describe, expect, it } from "vitest"; +import { SessionEventHub } from "../realtime/sessionEventHub.js"; +import { PiSessionService, type PiSessionManagerGateway } from "./piSessionService.js"; +import { registerSessionRoutes } from "./sessionRoutes.js"; + +let app: FastifyInstance; +let service: PiSessionService; +let sessionManager: RejectingSessionManager; + +beforeEach(async () => { + app = Fastify({ logger: false }); + await app.register(fastifyWebsocket); + sessionManager = new RejectingSessionManager(); + const eventHub = new SessionEventHub(); + service = new PiSessionService(eventHub, { sessionManager, heartbeatIntervalMs: 60_000 }); + registerSessionRoutes(app, service, eventHub); +}); + +afterEach(async () => { + await service.dispose(); + await app.close(); +}); + +describe("session routes", () => { + it("rejects prompt payloads that omit text without opening a session", async () => { + const response = await app.inject({ method: "POST", url: "/sessions/session-1/prompt", payload: { body: "Build the thing" } }); + + expect(response.statusCode).toBe(400); + expect(response.json()).toEqual({ error: "Prompt text is required" }); + expect(sessionManager.calls).toEqual({ create: 0, list: 0, listAll: 0, open: 0 }); + }); +}); + +class RejectingSessionManager implements PiSessionManagerGateway { + readonly calls = { create: 0, list: 0, listAll: 0, open: 0 }; + + list() { + this.calls.list += 1; + return Promise.resolve([]); + } + + create(): never { + this.calls.create += 1; + throw new Error("Session manager should not create sessions for invalid prompt payloads"); + } + + listAll() { + this.calls.listAll += 1; + return Promise.resolve([]); + } + + open(): never { + this.calls.open += 1; + throw new Error("Session manager should not open sessions for invalid prompt payloads"); + } +} diff --git a/src/server/sessions/sessionRoutes.ts b/src/server/sessions/sessionRoutes.ts index 4de5b55..5fa9aab 100644 --- a/src/server/sessions/sessionRoutes.ts +++ b/src/server/sessions/sessionRoutes.ts @@ -2,6 +2,11 @@ import type { FastifyInstance } from "fastify"; import type { SessionEventHub } from "../realtime/sessionEventHub.js"; import type { PiSessionService } from "./piSessionService.js"; +interface PromptRequestBody { + text?: unknown; + streamingBehavior?: unknown; +} + export function registerSessionRoutes(app: FastifyInstance, sessions: PiSessionService, eventHub: SessionEventHub, prefix = ""): void { app.get<{ Querystring: { cwd?: string } }>(`${prefix}/sessions`, async (request, reply) => { if (request.query.cwd === undefined || request.query.cwd === "") return reply.code(400).send({ error: "cwd query parameter is required" }); @@ -89,9 +94,9 @@ export function registerSessionRoutes(app: FastifyInstance, sessions: PiSessionS } }); - app.post<{ Params: { sessionId: string }; Body: { text: string; streamingBehavior?: "steer" | "followUp" } }>(`${prefix}/sessions/:sessionId/prompt`, async (request, reply) => { + app.post<{ Params: { sessionId: string }; Body: PromptRequestBody | undefined }>(`${prefix}/sessions/:sessionId/prompt`, async (request, reply) => { try { - await sessions.prompt(request.params.sessionId, request.body.text, request.body.streamingBehavior); + await sessions.prompt(request.params.sessionId, request.body?.text, request.body?.streamingBehavior); return { accepted: true }; } catch (error) { return reply.code(400).send({ error: error instanceof Error ? error.message : String(error) });