diff --git a/src/components/chat-message-list.tsx b/src/components/chat-message-list.tsx index f035689..c1b5724 100644 --- a/src/components/chat-message-list.tsx +++ b/src/components/chat-message-list.tsx @@ -2,7 +2,7 @@ import { type CSSProperties, type ReactElement, type ReactNode } from "react"; import { type VirtualItem } from "@tanstack/react-virtual"; -import { ImageIcon, MessageCircle } from "lucide-react"; +import { BarChart3, ImageIcon, MessageCircle } from "lucide-react"; import ReactMarkdown, { defaultUrlTransform, type Components, @@ -12,6 +12,7 @@ import remarkGfm from "remark-gfm"; import { ChatDiagramCard } from "@/components/chat-diagram-card"; import { useChatMessageListWorkflow } from "@/components/chat-message-list-workflow"; import { chatPanelModel } from "@/components/chat-panel-model"; +import { Button } from "@/components/ui/button"; import { ScrollArea } from "@/components/ui/scroll-area"; import { Spinner } from "@/components/ui/spinner"; import { @@ -86,6 +87,7 @@ const assistantMarkdownComponents: Components = { export type ChatMessageListProps = { readonly diagramStatesByMessageId?: Readonly>; readonly isDisabled?: boolean; + readonly isDiagramActionDisabled?: boolean; readonly isSending?: boolean; readonly messages?: readonly ChatMessageView[]; readonly needsLogin?: boolean; @@ -95,18 +97,21 @@ export type ChatMessageListProps = { ) => void; readonly pendingCitationId?: string | null; readonly pendingStatusText?: string | null; + readonly onCreateDiagram?: (message: ChatMessageView) => void; readonly sourceTitlesByDocumentId?: Readonly>; }; export function ChatMessageList({ diagramStatesByMessageId = {}, isDisabled = false, + isDiagramActionDisabled = false, isSending = false, messages = [], needsLogin = false, onCitationClick, pendingCitationId = null, pendingStatusText = null, + onCreateDiagram, sourceTitlesByDocumentId = {}, }: ChatMessageListProps): ReactElement { const { @@ -147,6 +152,8 @@ export function ChatMessageList({ diagramStatesByMessageId[getVirtualMessage(virtualItem)?.id ?? ""] } onCitationClick={onCitationClick} + isDiagramActionDisabled={isDiagramActionDisabled} + onCreateDiagram={onCreateDiagram} pendingCitationId={pendingCitationId} sourceTitlesByDocumentId={sourceTitlesByDocumentId} /> @@ -216,6 +223,8 @@ function VirtualMessageRow({ message, measureElement, onCitationClick, + isDiagramActionDisabled, + onCreateDiagram, pendingCitationId, sourceTitlesByDocumentId, }: { @@ -227,6 +236,8 @@ function VirtualMessageRow({ citation: ChatCitationView, citationId: string, ) => void; + readonly isDiagramActionDisabled: boolean; + readonly onCreateDiagram?: (message: ChatMessageView) => void; readonly pendingCitationId?: string | null; readonly sourceTitlesByDocumentId: Readonly>; }): ReactElement | null { @@ -249,8 +260,10 @@ function VirtualMessageRow({ > @@ -286,17 +299,21 @@ function EmptyChat({ function MessageBubble({ diagramState, + isDiagramActionDisabled, message, onCitationClick, + onCreateDiagram, pendingCitationId, sourceTitlesByDocumentId, }: { readonly diagramState: ChatDiagramState; + readonly isDiagramActionDisabled: boolean; readonly message: ChatMessageView; readonly onCitationClick?: ( citation: ChatCitationView, citationId: string, ) => void; + readonly onCreateDiagram?: (message: ChatMessageView) => void; readonly pendingCitationId?: string | null; readonly sourceTitlesByDocumentId: Readonly>; }): ReactElement { @@ -348,6 +365,12 @@ function MessageBubble({ )} + {displayImageCitations.length > 0 && (

@@ -381,6 +404,49 @@ function MessageBubble({ ); } +function AssistantDiagramAction({ + isDisabled, + message, + onCreateDiagram, + state, +}: { + readonly isDisabled: boolean; + readonly message: ChatMessageView; + readonly onCreateDiagram?: (message: ChatMessageView) => void; + readonly state: ChatDiagramState; +}): ReactElement | null { + if ( + !onCreateDiagram || + state.status === "loading" || + state.status === "ready" || + message.content.trim().length === 0 + ) { + return null; + } + + const label = + state.status === "empty" || state.status === "error" + ? "Try diagram again" + : "Create diagram"; + + return ( +

+ +
+ ); +} + function AssistantDiagram({ state, }: { @@ -403,7 +469,14 @@ function AssistantDiagram({ )} {state.status === "empty" && ( -

{state.reason}

+
+

+ No diagram created +

+

+ {state.reason} +

+
)} {state.status === "error" && (

{state.message}

diff --git a/src/components/chat-panel.test.ts b/src/components/chat-panel.test.ts index e0f6c02..1d9736e 100644 --- a/src/components/chat-panel.test.ts +++ b/src/components/chat-panel.test.ts @@ -123,6 +123,90 @@ describe("ChatPanel", () => { ).toBeTruthy(); }); + it("creates a diagram directly from an assistant answer", async () => { + const user = userEvent.setup(); + vi.mocked(workspaceClient.createChatDiagram).mockResolvedValue({ + diagram: { + type: "column", + source: "chart-visualization-skills", + title: "Revenue by Segment", + data: [ + { category: "Cloud", value: 42 }, + { category: "Ads", value: 28 }, + ], + }, + }); + + render( + React.createElement(C, { + messages: [ + { + id: "assistant_1", + role: "assistant", + content: "Cloud revenue was 42 and Ads revenue was 28.", + }, + ], + }), + ); + + await user.click( + screen.getByRole("button", { + name: "Create diagram for this answer", + }), + ); + + expect(workspaceClient.createChatDiagram).toHaveBeenCalledWith({ + answer: "Cloud revenue was 42 and Ads revenue was 28.", + }); + expect(await screen.findByText("Revenue by Segment")).toBeTruthy(); + expect( + screen.queryByRole("button", { + name: "Create diagram for this answer", + }), + ).toBeNull(); + }); + + it("shows a friendly no-diagram state for non-chartable answers", async () => { + const user = userEvent.setup(); + vi.mocked(workspaceClient.createChatDiagram).mockResolvedValue({ + diagram: { + type: "none", + reason: + "No clear chartable data was found. Ask for a table or numeric comparison first.", + }, + }); + + render( + React.createElement(C, { + messages: [ + { + id: "assistant_1", + role: "assistant", + content: "This is a qualitative summary without comparable numbers.", + }, + ], + }), + ); + + await user.click( + screen.getByRole("button", { + name: "Create diagram for this answer", + }), + ); + + expect(await screen.findByText("No diagram created")).toBeTruthy(); + expect( + screen.getByText( + "No clear chartable data was found. Ask for a table or numeric comparison first.", + ), + ).toBeTruthy(); + expect( + screen.getByRole("button", { + name: "Try diagram again for this answer", + }), + ).toBeTruthy(); + }); + it("treats slash diagram text as a local command instead of a chat message", async () => { const user = userEvent.setup(); const onSend = vi.fn(); diff --git a/src/components/chat-panel.tsx b/src/components/chat-panel.tsx index cf78d8d..1aeb897 100644 --- a/src/components/chat-panel.tsx +++ b/src/components/chat-panel.tsx @@ -112,10 +112,14 @@ export function ChatPanel({ !isSending && diagramTargetState?.status !== "loading"; - async function handleCreateDiagramCommand(): Promise { - if (!diagramTargetMessage || !canCreateDiagram) return; + async function handleCreateDiagramCommand( + targetMessage: ChatMessageView | undefined = diagramTargetMessage, + ): Promise { + if (!targetMessage || isDisabled || isSending) return; + + const messageId = targetMessage.id; + if (diagramStatesByMessageId[messageId]?.status === "loading") return; - const messageId = diagramTargetMessage.id; setDiagramStatesByMessageId((current) => ({ ...current, [messageId]: { status: "loading" }, @@ -123,7 +127,7 @@ export function ChatPanel({ try { const response = await workspaceClient.createChatDiagram({ - answer: diagramTargetMessage.content, + answer: targetMessage.content, }); setDiagramStatesByMessageId((current) => ({ ...current, @@ -254,21 +258,23 @@ export function ChatPanel({ handleCreateDiagramCommand()} onLoginClick={onLoginClick} onSend={handleComposerSend} /> diff --git a/src/domains/chat/diagram.test.ts b/src/domains/chat/diagram.test.ts index b699551..e2e6d50 100644 --- a/src/domains/chat/diagram.test.ts +++ b/src/domains/chat/diagram.test.ts @@ -1,16 +1,21 @@ -import { generateObject } from "ai" +import { generateObject, NoObjectGeneratedError } from "ai" import { afterEach, describe, expect, it, vi } from "vitest" import { + buildChatDiagramRepairPrompt, buildChatDiagramPrompt, generateChatDiagramSpec, parseChatDiagramRequestBody, retrieveAntvChartSkills, } from "./diagram" -vi.mock("ai", () => ({ - generateObject: vi.fn(), -})) +vi.mock("ai", async (importOriginal) => { + const actual = await importOriginal() + return { + ...actual, + generateObject: vi.fn(), + } +}) describe("parseChatDiagramRequestBody", () => { it("accepts trimmed answer content", () => { @@ -60,6 +65,21 @@ describe("buildChatDiagramPrompt", () => { }) }) +describe("buildChatDiagramRepairPrompt", () => { + it("guides failed object generation into a valid chart or no-diagram object", () => { + const prompt = buildChatDiagramRepairPrompt({ + answer: "Cloud revenue was 42 and Ads revenue was 28.", + failedOutput: "Here is a chart: ```json {} ```", + }) + + expect(prompt).toContain("Return one valid JSON object only") + expect(prompt).toContain("\"type\":\"none\"") + expect(prompt).toContain("\"source\":\"chart-visualization-skills\"") + expect(prompt).toContain("Previous invalid output") + expect(prompt).toContain("Cloud revenue was 42") + }) +}) + describe("generateChatDiagramSpec", () => { afterEach(() => { delete process.env.AI_GATEWAY_API_KEY @@ -102,6 +122,61 @@ describe("generateChatDiagramSpec", () => { }) }) + it("repairs schema-mismatched object output once before returning a chart", async () => { + process.env.AI_GATEWAY_API_KEY = "test_gateway_key" + vi.mocked(generateObject) + .mockRejectedValueOnce(makeNoObjectGeneratedError()) + .mockResolvedValueOnce({ + object: { + type: "column", + source: "chart-visualization-skills", + title: "Revenue by Segment", + data: [ + { category: "Cloud", value: 42 }, + { category: "Ads", value: 28 }, + ], + }, + } as Awaited>) + + const spec = await generateChatDiagramSpec({ + answer: "Cloud revenue was 42 and Ads revenue was 28.", + }) + + expect(generateObject).toHaveBeenCalledTimes(2) + expect(vi.mocked(generateObject).mock.calls[1]?.[0]).toEqual({ + model: "google/gemini-3-flash", + schema: expect.any(Object), + prompt: expect.stringContaining("The previous diagram-generation output"), + }) + expect(spec).toEqual({ + type: "column", + source: "chart-visualization-skills", + title: "Revenue by Segment", + axisXTitle: undefined, + axisYTitle: undefined, + data: [ + { category: "Cloud", time: undefined, value: 42 }, + { category: "Ads", time: undefined, value: 28 }, + ], + }) + }) + + it("returns a no-diagram response when schema repair also fails", async () => { + process.env.AI_GATEWAY_API_KEY = "test_gateway_key" + vi.mocked(generateObject).mockRejectedValue(makeNoObjectGeneratedError()) + + await expect( + generateChatDiagramSpec({ + answer: "This answer contains no chartable numbers.", + }), + ).resolves.toEqual({ + type: "none", + reason: + "No clear chartable data was found. Ask for a table or numeric comparison first.", + }) + expect(generateObject).toHaveBeenCalledTimes(2) + }) + it("normalizes sparse chart specs into no-diagram responses", async () => { process.env.AI_GATEWAY_API_KEY = "test_gateway_key" vi.mocked(generateObject).mockResolvedValue({ @@ -146,3 +221,31 @@ describe("generateChatDiagramSpec", () => { }) }) }) + +function makeNoObjectGeneratedError(): NoObjectGeneratedError { + return new NoObjectGeneratedError({ + message: "No object generated: response did not match schema.", + cause: new Error("schema mismatch"), + text: "Here is a chart that does not match the schema.", + response: { + id: "response_1", + modelId: "test-model", + timestamp: new Date("2026-01-01T00:00:00Z"), + }, + usage: { + inputTokens: 1, + inputTokenDetails: { + noCacheTokens: 1, + cacheReadTokens: 0, + cacheWriteTokens: 0, + }, + outputTokens: 1, + outputTokenDetails: { + textTokens: 1, + reasoningTokens: 0, + }, + totalTokens: 2, + }, + finishReason: "stop", + }) +} diff --git a/src/domains/chat/diagram.ts b/src/domains/chat/diagram.ts index f2d9327..41b29c8 100644 --- a/src/domains/chat/diagram.ts +++ b/src/domains/chat/diagram.ts @@ -1,4 +1,4 @@ -import { generateObject } from "ai" +import { generateObject, NoObjectGeneratedError } from "ai" import g2SkillIndex from "@antv/chart-visualization-skills/dist/index/g2.index.json" import type { Skill } from "@antv/chart-visualization-skills" import { z } from "zod" @@ -14,6 +14,9 @@ const ANTV_CHART_SKILL_TOP_K = 5 const MAX_ANTV_SKILL_QUERY_CHARS = 500 const MAX_ANTV_SKILL_CONTENT_CHARS = 2_400 const MAX_ANTV_SKILL_CONTEXT_CHARS = 16_000 +const MAX_FAILED_OBJECT_OUTPUT_CHARS = 4_000 +const DEFAULT_NO_DIAGRAM_REASON = + "No clear chartable data was found. Ask for a table or numeric comparison first." const ANTV_CHART_SEARCH_STOP_WORDS = new Set([ "the", "and", @@ -136,17 +139,69 @@ export async function generateChatDiagramSpec(input: { } const prompt = buildChatDiagramPrompt(input.answer) + let failedOutput: string | undefined + try { + return await requestChatDiagramObject({ + attempt: "initial", + prompt, + }) + } catch (error) { + if (!NoObjectGeneratedError.isInstance(error)) { + throw error + } + + failedOutput = error.text + logger.warn("chat-diagram: llm object generation failed", { + attempt: "initial", + detail: summarizeUnknownError(error), + generatedTextLength: error.text?.length ?? 0, + }) + } + + const repairPrompt = buildChatDiagramRepairPrompt({ + answer: input.answer, + failedOutput, + }) + + try { + return await requestChatDiagramObject({ + attempt: "repair", + prompt: repairPrompt, + }) + } catch (error) { + if (!NoObjectGeneratedError.isInstance(error)) { + throw error + } + + logger.warn("chat-diagram: llm object repair failed", { + attempt: "repair", + detail: summarizeUnknownError(error), + generatedTextLength: error.text?.length ?? 0, + }) + return { + type: "none", + reason: DEFAULT_NO_DIAGRAM_REASON, + } + } +} + +async function requestChatDiagramObject(input: { + readonly attempt: "initial" | "repair" + readonly prompt: string +}): Promise { logger.info("chat-diagram: llm request", { + attempt: input.attempt, model: CHAT_MODEL, - promptCharLength: prompt.length, + promptCharLength: input.prompt.length, }) const response = await generateObject({ model: CHAT_MODEL, schema: chatDiagramSpecSchema, - prompt, + prompt: input.prompt, }) const spec = normalizeChatDiagramSpec(response.object) logger.info("chat-diagram: llm response", { + attempt: input.attempt, model: CHAT_MODEL, type: spec.type, dataPointCount: spec.type === "none" ? 0 : spec.data.length, @@ -193,6 +248,43 @@ export function buildChatDiagramPrompt(answer: string): string { ].join("\n") } +export function buildChatDiagramRepairPrompt(input: { + readonly answer: string + readonly failedOutput?: string +}): string { + const failedOutput = input.failedOutput?.trim() + + return [ + "The previous diagram-generation output did not match the required schema.", + "Return one valid JSON object only. Do not include Markdown, prose, code fences, comments, or extra keys.", + "", + "Allowed JSON shapes:", + `{"type":"none","reason":"${DEFAULT_NO_DIAGRAM_REASON}"}`, + '{"type":"column","source":"chart-visualization-skills","title":"Revenue by Segment","axisYTitle":"USD billions","data":[{"category":"Cloud","value":42},{"category":"Ads","value":28}]}', + '{"type":"line","source":"chart-visualization-skills","title":"Revenue Trend","axisXTitle":"Period","axisYTitle":"USD billions","data":[{"time":"2024","value":10},{"time":"2025","value":14}]}', + "", + "Repair rules:", + "- If the answer does not contain at least two explicit comparable numbers or time points, return the none object.", + "- If making a chart would require inferred, fabricated, hidden, or unit-mixed values, return the none object.", + '- For charts, source must be exactly "chart-visualization-skills".', + "- For bar, column, and pie data, use category + value.", + "- For line data, use time + value.", + "- Use at most 12 data points.", + "", + "Answer content:", + input.answer, + failedOutput + ? [ + "", + "Previous invalid output:", + failedOutput.slice(0, MAX_FAILED_OBJECT_OUTPUT_CHARS), + ].join("\n") + : "", + ] + .filter((part): boolean => part.length > 0) + .join("\n") +} + export function getAntvChartSkillContext(answer: string): string { try { const skills = retrieveAntvChartSkills(