From 8b4d35798ac59a951ea29e39cdae1deb109e4c7b Mon Sep 17 00:00:00 2001 From: Mathias Fredriksson Date: Wed, 18 Mar 2026 15:50:57 +0200 Subject: [PATCH] refactor: type both chat message parsers (#23176) Both message parsers accepted untyped input and relied on scattered asRecord/asString calls to extract fields at runtime. With the discriminated ChatMessagePart union, both accept typed input directly and narrow via switch (part.type). parseMessageContent narrows from (content: unknown) to (content: readonly ChatMessagePart[] | undefined), removing legacy input shape handling the Go backend normalizes away. applyMessagePartToStreamState narrows from Record to ChatMessagePart. The SSE type guards had a & Record intersection that widened everything untyped downstream. Since the data comes from our own API, the intersection was removed and all handlers in ChatContext now use generated types directly. Fixes tool_call_id and tool_name variant tags in codersdk/chats.go: marked optional to match reality (Go guards against empty values, omitempty omits them at the wire level). Refs #23168, #23175 --- codersdk/chats.go | 4 +- codersdk/chats_test.go | 11 - site/src/api/typesGenerated.ts | 8 +- .../AgentsPage/AgentDetail/ChatContext.ts | 66 ++--- .../AgentDetail/messageParsing.test.ts | 65 +---- .../AgentsPage/AgentDetail/messageParsing.ts | 247 +++++++----------- .../AgentDetail/streamState.test.ts | 70 +++-- .../AgentsPage/AgentDetail/streamState.ts | 59 ++--- 8 files changed, 207 insertions(+), 323 deletions(-) diff --git a/codersdk/chats.go b/codersdk/chats.go index 3255cf3d05..403050c4c0 100644 --- a/codersdk/chats.go +++ b/codersdk/chats.go @@ -129,8 +129,8 @@ type ChatMessagePart struct { Type ChatMessagePartType `json:"type"` Text string `json:"text,omitempty" variants:"text,reasoning"` Signature string `json:"signature,omitempty"` - ToolCallID string `json:"tool_call_id,omitempty" variants:"tool-call,tool-result"` - ToolName string `json:"tool_name,omitempty" variants:"tool-call,tool-result"` + ToolCallID string `json:"tool_call_id,omitempty" variants:"tool-call?,tool-result?"` + ToolName string `json:"tool_name,omitempty" variants:"tool-call?,tool-result?"` Args json.RawMessage `json:"args,omitempty" variants:"tool-call?"` ArgsDelta string `json:"args_delta,omitempty" variants:"tool-call?"` Result json.RawMessage `json:"result,omitempty" variants:"tool-result?"` diff --git a/codersdk/chats_test.go b/codersdk/chats_test.go index b2c3246bdf..a698b14505 100644 --- a/codersdk/chats_test.go +++ b/codersdk/chats_test.go @@ -223,7 +223,6 @@ func TestChatMessagePartVariantTags(t *testing.T) { // Parse all variants tags from the struct and validate them. typ := reflect.TypeOf(codersdk.ChatMessagePart{}) coveredTypes := make(map[codersdk.ChatMessagePartType]bool) - hasRequired := make(map[codersdk.ChatMessagePartType]bool) for i := range typ.NumField() { f := typ.Field(i) @@ -245,7 +244,6 @@ func TestChatMessagePartVariantTags(t *testing.T) { "the discriminant field must not have a variants tag; %s", editHint) for _, entry := range strings.Split(varTag, ",") { - isOptional := strings.HasSuffix(entry, "?") typeLit := codersdk.ChatMessagePartType(strings.TrimSuffix(entry, "?")) assert.True(t, knownTypes[typeLit], @@ -253,9 +251,6 @@ func TestChatMessagePartVariantTags(t *testing.T) { f.Name, typeLit, editHint) coveredTypes[typeLit] = true - if !isOptional { - hasRequired[typeLit] = true - } } } @@ -264,12 +259,6 @@ func TestChatMessagePartVariantTags(t *testing.T) { assert.True(t, coveredTypes[pt], "ChatMessagePartType %q is not referenced by any variants tag; %s", pt, editHint) } - - // Every variant must have at least one required field. - for pt := range coveredTypes { - assert.True(t, hasRequired[pt], - "variant %q has no required fields (all have ? suffix); %s", pt, editHint) - } } func TestModelCostConfig_LegacyNumericJSON(t *testing.T) { diff --git a/site/src/api/typesGenerated.ts b/site/src/api/typesGenerated.ts index bdc2e128b6..46113d0a99 100644 --- a/site/src/api/typesGenerated.ts +++ b/site/src/api/typesGenerated.ts @@ -1805,8 +1805,8 @@ export interface ChatTextPart { // From codersdk/chats.go export interface ChatToolCallPart { readonly type: "tool-call"; - readonly tool_call_id: string; - readonly tool_name: string; + readonly tool_call_id?: string; + readonly tool_name?: string; readonly args?: Record; readonly args_delta?: string; /** @@ -1819,8 +1819,8 @@ export interface ChatToolCallPart { // From codersdk/chats.go export interface ChatToolResultPart { readonly type: "tool-result"; - readonly tool_call_id: string; - readonly tool_name: string; + readonly tool_call_id?: string; + readonly tool_name?: string; readonly result?: Record; readonly is_error?: boolean; /** diff --git a/site/src/pages/AgentsPage/AgentDetail/ChatContext.ts b/site/src/pages/AgentsPage/AgentDetail/ChatContext.ts index 65d89ee185..4be26aabed 100644 --- a/site/src/pages/AgentsPage/AgentDetail/ChatContext.ts +++ b/site/src/pages/AgentsPage/AgentDetail/ChatContext.ts @@ -1,7 +1,7 @@ import { watchChat } from "api/api"; import { chatMessagesKey, updateInfiniteChatsCache } from "api/queries/chats"; import type * as TypesGen from "api/typesGenerated"; -import { asRecord, asString } from "components/ai-elements/runtimeTypeUtils"; + import { startTransition, useCallback, @@ -16,21 +16,7 @@ import type { ChatDetailError } from "../usageLimitMessage"; import { applyMessagePartToStreamState } from "./streamState"; import type { StreamState } from "./types"; -const VALID_CHAT_STATUSES: ReadonlySet = new Set([ - "pending", - "running", - "completed", - "error", - "paused", - "waiting", -]); - -const isValidChatStatus = (value: unknown): value is TypesGen.ChatStatus => - typeof value === "string" && VALID_CHAT_STATUSES.has(value); - -const isChatStreamEvent = ( - data: unknown, -): data is TypesGen.ChatStreamEvent & Record => +const isChatStreamEvent = (data: unknown): data is TypesGen.ChatStreamEvent => typeof data === "object" && data !== null && "type" in data && @@ -38,12 +24,10 @@ const isChatStreamEvent = ( const isChatStreamEventArray = ( data: unknown, -): data is (TypesGen.ChatStreamEvent & Record)[] => +): data is TypesGen.ChatStreamEvent[] => Array.isArray(data) && data.every(isChatStreamEvent); -const toChatStreamEvents = ( - data: unknown, -): (TypesGen.ChatStreamEvent & Record)[] => { +const toChatStreamEvents = (data: unknown): TypesGen.ChatStreamEvent[] => { if (isChatStreamEvent(data)) { return [data]; } @@ -158,8 +142,8 @@ type ChatStore = { isDuplicate: boolean; changed: boolean; }; - applyMessagePart: (part: Record) => void; - applyMessageParts: (parts: readonly Record[]) => void; + applyMessagePart: (part: TypesGen.ChatMessagePart) => void; + applyMessageParts: (parts: readonly TypesGen.ChatMessagePart[]) => void; setQueuedMessages: ( queuedMessages: readonly TypesGen.ChatQueuedMessage[] | undefined, ) => void; @@ -281,7 +265,7 @@ export const createChatStore = (): ChatStore => { return { isDuplicate, changed: actuallyChanged }; }; - const applyMessageParts = (parts: readonly Record[]) => { + const applyMessageParts = (parts: readonly TypesGen.ChatMessagePart[]) => { if (parts.length === 0) { return; } @@ -657,7 +641,7 @@ export const useChatStore = ( return currentStatus !== "pending" && currentStatus !== "waiting"; }; - const pendingMessageParts: Record[] = []; + const pendingMessageParts: TypesGen.ChatMessagePart[] = []; const flushMessageParts = () => { if (pendingMessageParts.length === 0) { return; @@ -681,14 +665,13 @@ export const useChatStore = ( for (const streamEvent of streamEvents) { if (streamEvent.type === "message_part") { - const eventChatID = asString(streamEvent.chat_id); - if (eventChatID && eventChatID !== chatID) { + if (streamEvent.chat_id && streamEvent.chat_id !== chatID) { continue; } if (!shouldApplyMessagePart()) { continue; } - const part = asRecord(streamEvent.message_part?.part); + const part = streamEvent.message_part?.part; if (part) { cancelScheduledStreamReset(); pendingMessageParts.push(part); @@ -703,8 +686,7 @@ export const useChatStore = ( if (!message) { continue; } - const eventChatID = asString(streamEvent.chat_id); - if (eventChatID && eventChatID !== chatID) { + if (streamEvent.chat_id && streamEvent.chat_id !== chatID) { continue; } const { changed } = store.upsertDurableMessage(message); @@ -730,26 +712,21 @@ export const useChatStore = ( continue; } case "queue_update": - { - const eventChatID = asString(streamEvent.chat_id); - if (eventChatID && eventChatID !== chatID) { - continue; - } + if (streamEvent.chat_id && streamEvent.chat_id !== chatID) { + continue; } wsQueueUpdateReceivedRef.current = true; store.setQueuedMessages(streamEvent.queued_messages); updateChatQueuedMessages(streamEvent.queued_messages); continue; case "status": { - const status = asRecord(streamEvent.status); - const nextStatus = asString(status?.status); - if (!isValidChatStatus(nextStatus)) { + const nextStatus = streamEvent.status?.status; + if (!nextStatus) { continue; } - const eventChatID = asString(streamEvent.chat_id); - if (eventChatID && eventChatID !== chatID) { - store.setSubagentStatusOverride(eventChatID, nextStatus); + if (streamEvent.chat_id && streamEvent.chat_id !== chatID) { + store.setSubagentStatusOverride(streamEvent.chat_id, nextStatus); continue; } @@ -771,13 +748,11 @@ export const useChatStore = ( continue; } case "error": { - const eventChatID = asString(streamEvent.chat_id); - if (eventChatID && eventChatID !== chatID) { + if (streamEvent.chat_id && streamEvent.chat_id !== chatID) { continue; } - const error = asRecord(streamEvent.error); const reason = - asString(error?.message).trim() || "Chat processing failed."; + streamEvent.error?.message.trim() || "Chat processing failed."; store.setChatStatus("error"); store.setStreamError(reason); store.clearRetryState(); @@ -792,8 +767,7 @@ export const useChatStore = ( continue; } case "retry": { - const eventChatID = asString(streamEvent.chat_id); - if (eventChatID && eventChatID !== chatID) { + if (streamEvent.chat_id && streamEvent.chat_id !== chatID) { continue; } const retry = streamEvent.retry; diff --git a/site/src/pages/AgentsPage/AgentDetail/messageParsing.test.ts b/site/src/pages/AgentsPage/AgentDetail/messageParsing.test.ts index 0adebe26ac..0a2e1de5d7 100644 --- a/site/src/pages/AgentsPage/AgentDetail/messageParsing.test.ts +++ b/site/src/pages/AgentsPage/AgentDetail/messageParsing.test.ts @@ -49,14 +49,6 @@ describe("parseToolResultIsError", () => { }); describe("parseMessageContent", () => { - it("returns empty result for null content", () => { - const result = parseMessageContent(null); - expect(result.markdown).toBe(""); - expect(result.blocks).toEqual([]); - expect(result.toolCalls).toEqual([]); - expect(result.toolResults).toEqual([]); - }); - it("returns empty result for undefined content", () => { const result = parseMessageContent(undefined); expect(result.markdown).toBe(""); @@ -71,12 +63,6 @@ describe("parseMessageContent", () => { expect(result.toolResults).toEqual([]); }); - it("handles a plain string content", () => { - const result = parseMessageContent("Hello world"); - expect(result.markdown).toBe("Hello world"); - expect(result.blocks).toEqual([]); - }); - it("parses a single text block", () => { const result = parseMessageContent([{ type: "text", text: "Hello" }]); expect(result.markdown).toBe("Hello"); @@ -173,7 +159,7 @@ describe("parseMessageContent", () => { type: "tool-result", tool_name: "bash", tool_call_id: "call-1", - result: "ok", + result: { output: "ok" }, }, { type: "text", text: "Done!" }, ]); @@ -195,32 +181,6 @@ describe("parseMessageContent", () => { expect(result.toolCalls[0].id).toBe("tool-call-0"); }); - it("handles unknown block types gracefully (no crash)", () => { - const result = parseMessageContent([ - { type: "unknown_block_type", text: "some text" }, - ]); - // Unknown types fall through to the default branch which treats - // the text field as a response. - expect(result.markdown).toBe("some text"); - expect(result.blocks).toEqual([{ type: "response", text: "some text" }]); - }); - - it("handles non-object array entries gracefully", () => { - const result = parseMessageContent(["raw string", 42, null]); - expect(result.markdown).toBe("raw string"); - expect(result.blocks).toEqual([{ type: "response", text: "raw string" }]); - }); - - it("handles an object with a type field (treated as single-element array)", () => { - const result = parseMessageContent({ type: "text", text: "single" }); - expect(result.markdown).toBe("single"); - }); - - it("handles an object with text/content fields", () => { - const result = parseMessageContent({ text: "fallback text" }); - expect(result.markdown).toBe("fallback text"); - }); - it("extracts fileId from a file block with file_id", () => { const result = parseMessageContent([ { @@ -255,6 +215,16 @@ describe("parseMessageContent", () => { }); }); + it("skips file parts without data or file_id", () => { + const result = parseMessageContent([ + { + type: "file", + media_type: "image/png", + }, + ]); + expect(result.blocks).toHaveLength(0); + }); + it("parses a file-reference block into blocks", () => { const result = parseMessageContent([ { @@ -275,19 +245,6 @@ describe("parseMessageContent", () => { }); }); - it("defaults lines to 0 when no line fields are provided", () => { - const result = parseMessageContent([ - { - type: "file-reference", - file_name: "bare.ts", - content: "bare content", - }, - ]); - const ref = result.blocks[0] as { start_line: number; end_line: number }; - expect(ref.start_line).toBe(0); - expect(ref.end_line).toBe(0); - }); - it("does not affect markdown when file-reference blocks are present", () => { const result = parseMessageContent([ { type: "text", text: "Hello" }, diff --git a/site/src/pages/AgentsPage/AgentDetail/messageParsing.ts b/site/src/pages/AgentsPage/AgentDetail/messageParsing.ts index c65bd5cb31..15e2dd5e8f 100644 --- a/site/src/pages/AgentsPage/AgentDetail/messageParsing.ts +++ b/site/src/pages/AgentsPage/AgentDetail/messageParsing.ts @@ -115,166 +115,101 @@ export const mergeTools = ( return merged; }; -export const parseMessageContent = (content: unknown): ParsedMessageContent => { - if (typeof content === "string") { - return { - ...emptyParsedMessageContent(), - markdown: content, - }; - } - - if (Array.isArray(content)) { - const parsed = emptyParsedMessageContent(); - for (const [index, block] of content.entries()) { - if (typeof block === "string") { - parsed.markdown = appendText(parsed.markdown, block); - parsed.blocks = appendTextBlock(parsed.blocks, "response", block); - continue; - } - - const typedBlock = asRecord(block); - if (!typedBlock) { - continue; - } - - switch (asString(typedBlock.type)) { - case "text": { - const text = asString(typedBlock.text); - parsed.markdown = appendText(parsed.markdown, text); - parsed.blocks = appendTextBlock(parsed.blocks, "response", text); - break; - } - case "reasoning": { - const text = asString(typedBlock.text); - parsed.reasoning = appendText(parsed.reasoning, text); - parsed.blocks = appendTextBlock(parsed.blocks, "thinking", text); - break; - } - case "tool-call": { - // Provider-executed tool calls (e.g. web_search) are - // handled by the provider itself — hide them from the - // tool card UI and let the sources component render - // their results. - if (typedBlock.provider_executed) { - break; - } - const name = asString(typedBlock.tool_name); - const id = asString(typedBlock.tool_call_id) || `tool-call-${index}`; - parsed.toolCalls.push({ - id, - name: name || "Tool", - args: typedBlock.args, - }); - parsed.blocks = ensureToolBlock(parsed.blocks, id); - break; - } - case "file-reference": { - const fileName = asString(typedBlock.file_name); - const startLine = Number(typedBlock.start_line) || 0; - const endLine = Number(typedBlock.end_line) || startLine; - const content = asString(typedBlock.content); - parsed.blocks.push({ - type: "file-reference", - file_name: fileName, - start_line: startLine, - end_line: endLine, - content, - }); - break; - } - case "tool-result": { - // Skip synthetic results for provider-executed tools. - if (typedBlock.provider_executed) { - break; - } - const name = asString(typedBlock.tool_name); - const id = - asString(typedBlock.tool_call_id) || `tool-result-${index}`; - const result = typedBlock.result; - parsed.toolResults.push({ - id, - name: name || "Tool", - result, - isError: parseToolResultIsError(name, typedBlock, result), - }); - parsed.blocks = ensureToolBlock(parsed.blocks, id); - break; - } - case "file": { - const mediaType = asString(typedBlock.media_type); - const data = asString(typedBlock.data) || undefined; - const fileId = asString(typedBlock.file_id) || undefined; - if (mediaType && (data || fileId)) { - parsed.blocks = [ - ...parsed.blocks, - { type: "file", media_type: mediaType, data, file_id: fileId }, - ]; - } - break; - } - case "source": { - const url = asString(typedBlock.url); - const title = asString(typedBlock.title); - if (url) { - const source = { url, title: title || url }; - // Still populate the flat list for backward compat. - if (!parsed.sources.some((s) => s.url === url)) { - parsed.sources.push(source); - } - // Group consecutive sources into a single - // inline block at this position. - const lastBlock = parsed.blocks[parsed.blocks.length - 1]; - if ( - lastBlock && - lastBlock.type === "sources" && - !lastBlock.sources.some((s) => s.url === url) - ) { - lastBlock.sources.push(source); - } else if (!lastBlock || lastBlock.type !== "sources") { - parsed.blocks.push({ - type: "sources", - sources: [source], - }); - } - } - break; - } - default: { - const text = asString(typedBlock.text); - parsed.markdown = appendText(parsed.markdown, text); - parsed.blocks = appendTextBlock(parsed.blocks, "response", text); - break; - } - } - } - return parsed; - } - - if (content === null || content === undefined) { +export const parseMessageContent = ( + content: readonly TypesGen.ChatMessagePart[] | undefined, +): ParsedMessageContent => { + if (!content || content.length === 0) { return emptyParsedMessageContent(); } - const typedContent = asRecord(content); - if (!typedContent) { - const markdown = String(content); - return { - ...emptyParsedMessageContent(), - markdown, - blocks: appendTextBlock([], "response", markdown), - }; + const parsed = emptyParsedMessageContent(); + for (const [index, part] of content.entries()) { + switch (part.type) { + case "text": { + parsed.markdown = appendText(parsed.markdown, part.text); + parsed.blocks = appendTextBlock(parsed.blocks, "response", part.text); + break; + } + case "reasoning": { + parsed.reasoning = appendText(parsed.reasoning, part.text); + parsed.blocks = appendTextBlock(parsed.blocks, "thinking", part.text); + break; + } + case "tool-call": { + // Provider-executed tool calls (e.g. web_search) are + // handled by the provider itself — hide them from the + // tool card UI and let the sources component render + // their results. + if (part.provider_executed) { + break; + } + const id = part.tool_call_id || `tool-call-${index}`; + parsed.toolCalls.push({ + id, + name: part.tool_name || "Tool", + args: part.args, + }); + parsed.blocks = ensureToolBlock(parsed.blocks, id); + break; + } + case "file-reference": { + parsed.blocks.push(part); + break; + } + case "tool-result": { + // Skip synthetic results for provider-executed tools. + if (part.provider_executed) { + break; + } + const id = part.tool_call_id || `tool-result-${index}`; + const name = part.tool_name || "Tool"; + parsed.toolResults.push({ + id, + name, + result: part.result, + isError: parseToolResultIsError(name, part, part.result), + }); + parsed.blocks = ensureToolBlock(parsed.blocks, id); + break; + } + case "file": { + if (part.data || part.file_id) { + parsed.blocks = [...parsed.blocks, part]; + } + break; + } + case "source": { + if (part.url) { + const source = { url: part.url, title: part.title || part.url }; + // Still populate the flat list for backward compat. + if (!parsed.sources.some((s) => s.url === part.url)) { + parsed.sources.push(source); + } + // Group consecutive sources into a single + // inline block at this position. + const lastBlock = parsed.blocks[parsed.blocks.length - 1]; + if ( + lastBlock && + lastBlock.type === "sources" && + !lastBlock.sources.some((s) => s.url === part.url) + ) { + lastBlock.sources.push(source); + } else if (!lastBlock || lastBlock.type !== "sources") { + parsed.blocks.push({ + type: "sources", + sources: [source], + }); + } + } + break; + } + default: { + const _exhaustive: never = part; + break; + } + } } - - if (typedContent.type) { - return parseMessageContent([typedContent]); - } - - const markdown = - asString(typedContent.text) || asString(typedContent.content); - return { - ...emptyParsedMessageContent(), - markdown, - blocks: appendTextBlock([], "response", markdown), - }; + return parsed; }; export const parseMessagesWithMergedTools = ( diff --git a/site/src/pages/AgentsPage/AgentDetail/streamState.test.ts b/site/src/pages/AgentsPage/AgentDetail/streamState.test.ts index 2f3230240b..3918290f44 100644 --- a/site/src/pages/AgentsPage/AgentDetail/streamState.test.ts +++ b/site/src/pages/AgentsPage/AgentDetail/streamState.test.ts @@ -139,13 +139,13 @@ describe("applyMessagePartToStreamState", () => { const callIds = Object.keys(state!.toolCalls); expect(callIds).toHaveLength(2); - // First result arrives without an explicit tool_call_id. + // First result arrives without a tool_call_id. state = applyMessagePartToStreamState(state, { type: "tool-result", tool_name: "bash", result: { output: "file.txt" }, }); - // Second result arrives without an explicit tool_call_id. + // Second result arrives without a tool_call_id. state = applyMessagePartToStreamState(state, { type: "tool-result", tool_name: "bash", @@ -179,21 +179,6 @@ describe("applyMessagePartToStreamState", () => { }); }); - it("returns prev for unknown part type", () => { - const prev = createEmptyStreamState(); - const result = applyMessagePartToStreamState(prev, { - type: "banana", - }); - expect(result).toBe(prev); - }); - - it("returns null for unknown part type when prev is null", () => { - const result = applyMessagePartToStreamState(null, { - type: "banana", - }); - expect(result).toBeNull(); - }); - it("accumulates multiple tool calls in sequence", () => { let state: StreamState | null = null; state = applyMessagePartToStreamState(state, { @@ -268,6 +253,57 @@ describe("applyMessagePartToStreamState", () => { expect(prev.toolResults).toEqual({}); }); + it("adds a file block from a file part with data", () => { + const result = applyMessagePartToStreamState(null, { + type: "file", + media_type: "image/png", + data: "iVBORw0KGgo=", + }); + expect(result).not.toBeNull(); + expect(result!.blocks).toHaveLength(1); + expect(result!.blocks[0]).toMatchObject({ + type: "file", + media_type: "image/png", + data: "iVBORw0KGgo=", + }); + }); + + it("adds a file block from a file part with file_id", () => { + const result = applyMessagePartToStreamState(null, { + type: "file", + media_type: "image/png", + file_id: "abc-123", + }); + expect(result).not.toBeNull(); + expect(result!.blocks).toHaveLength(1); + expect(result!.blocks[0]).toMatchObject({ + type: "file", + media_type: "image/png", + file_id: "abc-123", + }); + }); + + it("returns prev for file part without data or file_id", () => { + const prev = createEmptyStreamState(); + const result = applyMessagePartToStreamState(prev, { + type: "file", + media_type: "image/png", + }); + expect(result).toBe(prev); + }); + + it("returns prev for file-reference part (not a streaming type)", () => { + const prev = createEmptyStreamState(); + const result = applyMessagePartToStreamState(prev, { + type: "file-reference", + file_name: "main.go", + start_line: 1, + end_line: 10, + content: "package main", + }); + expect(result).toBe(prev); + }); + it("adds a sources block from a source part", () => { let state: StreamState | null = null; state = applyMessagePartToStreamState(state, { diff --git a/site/src/pages/AgentsPage/AgentDetail/streamState.ts b/site/src/pages/AgentsPage/AgentDetail/streamState.ts index 3e04f20b0d..9272699cf2 100644 --- a/site/src/pages/AgentsPage/AgentDetail/streamState.ts +++ b/site/src/pages/AgentsPage/AgentDetail/streamState.ts @@ -1,4 +1,4 @@ -import { asString } from "components/ai-elements/runtimeTypeUtils"; +import type * as TypesGen from "api/typesGenerated"; import { appendTextBlock } from "./blockUtils"; import { ensureToolBlock, parseToolResultIsError } from "./messageParsing"; import { mergeStreamPayload } from "./streamingJson"; @@ -15,30 +15,27 @@ export const createEmptyStreamState = (): StreamState => ({ export const applyMessagePartToStreamState = ( prev: StreamState | null, - part: Record, + part: TypesGen.ChatMessagePart, ): StreamState | null => { - const partType = asString(part.type); const nextState: StreamState = prev ?? createEmptyStreamState(); - switch (partType) { + switch (part.type) { case "text": { - const text = asString(part.text); - if (!text) { + if (!part.text) { return prev; } return { ...nextState, - blocks: appendTextBlock(nextState.blocks, "response", text), + blocks: appendTextBlock(nextState.blocks, "response", part.text), }; } case "reasoning": { - const text = asString(part.text); - if (!text) { + if (!part.text) { return prev; } return { ...nextState, - blocks: appendTextBlock(nextState.blocks, "thinking", text), + blocks: appendTextBlock(nextState.blocks, "thinking", part.text), }; } case "tool-call": { @@ -48,12 +45,11 @@ export const applyMessagePartToStreamState = ( if (part.provider_executed) { return prev; } - const toolName = asString(part.tool_name); const existingByName = Object.values(nextState.toolCalls).find( - (call) => call.name === toolName, + (call) => call.name === part.tool_name, ); const toolCallID = - asString(part.tool_call_id) || + part.tool_call_id || (existingByName && !existingByName.args ? existingByName.id : null) || `tool-call-${Object.keys(nextState.toolCalls).length + 1}-${++nextFallbackID}`; const existing = nextState.toolCalls[toolCallID]; @@ -71,7 +67,7 @@ export const applyMessagePartToStreamState = ( ...nextState.toolCalls, [toolCallID]: { id: toolCallID, - name: toolName || existing?.name || "Tool", + name: part.tool_name || existing?.name || "Tool", args: nextArgs.value, argsRaw: nextArgs.rawText, }, @@ -83,15 +79,14 @@ export const applyMessagePartToStreamState = ( if (part.provider_executed) { return prev; } - const toolName = asString(part.tool_name); const existingByName = Object.values(nextState.toolResults).find( - (result) => result.name === toolName, + (result) => result.name === part.tool_name, ); const existingCallByName = Object.values(nextState.toolCalls).find( - (call) => call.name === toolName, + (call) => call.name === part.tool_name, ); const toolCallID = - asString(part.tool_call_id) || + part.tool_call_id || (existingByName && !existingByName.result ? existingByName.id : null) || (existingCallByName && !nextState.toolResults[existingCallByName.id] ? existingCallByName.id @@ -104,7 +99,7 @@ export const applyMessagePartToStreamState = ( part.result, undefined, // no delta: tool results arrive complete, not streamed incrementally ); - const nextToolName = toolName || existing?.name || "Tool"; + const nextToolName = part.tool_name || existing?.name || "Tool"; const nextIsError = existing?.isError || parseToolResultIsError(nextToolName, part, nextResult.value); @@ -125,29 +120,21 @@ export const applyMessagePartToStreamState = ( }; } case "file": { - const mediaType = asString(part.media_type); - const data = asString(part.data) || undefined; - const fileId = asString(part.file_id) || undefined; - if (!mediaType || (!data && !fileId)) { + if (!part.data && !part.file_id) { return prev; } return { ...nextState, - blocks: [ - ...nextState.blocks, - { type: "file", media_type: mediaType, data, file_id: fileId }, - ], + blocks: [...nextState.blocks, part], }; } case "source": { - const url = asString(part.url); - const title = asString(part.title); - if (!url) { + if (!part.url) { return prev; } - const source = { url, title: title || url }; + const source = { url: part.url, title: part.title || part.url }; // Still populate the flat list for backward compat. - if (nextState.sources.some((s) => s.url === url)) { + if (nextState.sources.some((s) => s.url === part.url)) { return prev; } const newSources = [...nextState.sources, source]; @@ -174,8 +161,14 @@ export const applyMessagePartToStreamState = ( blocks: newBlocks, }; } - default: + // file-reference parts only appear in persisted messages + // from user input, never via SSE streaming. + case "file-reference": return prev; + default: { + const _exhaustive: never = part; + return prev; + } } };