mirror of
https://github.com/t8y2/dbx.git
synced 2026-10-02 02:34:42 +08:00
fix(ai): keep replies visible on legacy WebView
This commit is contained in:
@@ -80,6 +80,26 @@ describe("createAiMessageRenderer", () => {
|
||||
expect(code).toMatchObject({ type: "code", pending: false, html: "<span>SELECT 1</span>" });
|
||||
});
|
||||
|
||||
it("falls back to escaped code when highlighting is not supported", () => {
|
||||
const markdown = (text: string) => `<p>${text}</p>`;
|
||||
const highlightCode = vi.fn(() => {
|
||||
throw new SyntaxError("Invalid regular expression: invalid group specifier name");
|
||||
});
|
||||
const renderer = createAiMessageRenderer({ markdown, highlightCode });
|
||||
|
||||
const [code] = renderer.render("```sql\nSELECT < 1\n```");
|
||||
|
||||
expect(highlightCode).toHaveBeenCalledWith("SELECT < 1", "SQL");
|
||||
expect(code).toEqual({
|
||||
type: "code",
|
||||
content: "SELECT < 1",
|
||||
html: "SELECT < 1",
|
||||
lang: "SQL",
|
||||
isSql: true,
|
||||
pending: false,
|
||||
});
|
||||
});
|
||||
|
||||
it("re-parses only the last paragraph of a long streaming answer", () => {
|
||||
const markdown = vi.fn((text: string) => `<p>${text}</p>`);
|
||||
const renderer = createAiMessageRenderer({ markdown });
|
||||
|
||||
@@ -97,7 +97,7 @@ export function createAiMessageRenderer(options: AiMessageRendererOptions) {
|
||||
}
|
||||
const lang = normalizeAiCodeLanguage(segment.lang);
|
||||
// Highlighting a block that is still streaming is wasted work: it is re-highlighted once the fence closes.
|
||||
const highlighted = flags.live && flags.pending ? undefined : options.highlightCode?.(segment.content, lang);
|
||||
const highlighted = flags.live && flags.pending ? undefined : safeHighlightCode(segment.content, lang);
|
||||
return {
|
||||
type: "code",
|
||||
content: segment.content,
|
||||
@@ -108,6 +108,15 @@ export function createAiMessageRenderer(options: AiMessageRendererOptions) {
|
||||
};
|
||||
}
|
||||
|
||||
function safeHighlightCode(content: string, lang: string): string | undefined {
|
||||
if (!options.highlightCode) return undefined;
|
||||
try {
|
||||
return options.highlightCode(content, lang);
|
||||
} catch {
|
||||
return undefined;
|
||||
}
|
||||
}
|
||||
|
||||
function renderCachedSegment(segment: MessageSegment, flags: SegmentRenderFlags): AiMessageRenderSegment {
|
||||
if (segment.content.length > maxCacheableChars) return renderSegment(segment, flags);
|
||||
|
||||
|
||||
@@ -441,6 +441,10 @@ export type AgentEvent =
|
||||
}
|
||||
| { type: "error"; message: string };
|
||||
|
||||
type TauriAgentEvent = AgentEvent & {
|
||||
session_id?: string;
|
||||
};
|
||||
|
||||
export async function aiAgentStream(
|
||||
sessionId: string,
|
||||
request: AiCompletionRequest,
|
||||
@@ -457,9 +461,11 @@ export async function aiAgentStream(
|
||||
confirmedSchema?: string,
|
||||
_signal?: AbortSignal,
|
||||
): Promise<string> {
|
||||
const unlisten: UnlistenFn = await listen<AgentEvent>("ai-agent-event", (event) => {
|
||||
onEvent(event.payload);
|
||||
if (event.payload.type === "agent_end" || event.payload.type === "error") {
|
||||
const unlisten: UnlistenFn = await listen<TauriAgentEvent>("ai-agent-event", (event) => {
|
||||
const payload = event.payload;
|
||||
if (payload.session_id && payload.session_id !== sessionId) return;
|
||||
onEvent(payload);
|
||||
if (payload.type === "agent_end" || payload.type === "error") {
|
||||
unlisten();
|
||||
}
|
||||
});
|
||||
|
||||
+22
@@ -0,0 +1,22 @@
|
||||
import assert from "node:assert/strict";
|
||||
import { readFileSync } from "node:fs";
|
||||
import { test } from "vitest";
|
||||
|
||||
const tauriBackendSource = readFileSync("apps/desktop/src/lib/backend/tauri.ts", "utf8");
|
||||
const tauriAiCommandSource = readFileSync("src-tauri/src/commands/ai.rs", "utf8");
|
||||
|
||||
test("Tauri AI agent stream events carry their session id", () => {
|
||||
assert.match(tauriAiCommandSource, /struct AiAgentEventPayload \{/);
|
||||
assert.match(tauriAiCommandSource, /session_id: String/);
|
||||
assert.match(tauriAiCommandSource, /#\[serde\(flatten\)\]\s*event: AgentEvent/);
|
||||
assert.match(tauriAiCommandSource, /let event_session_id = session_id\.clone\(\);/);
|
||||
assert.match(tauriAiCommandSource, /AiAgentEventPayload \{ session_id: event_session_id\.clone\(\), event \}/);
|
||||
assert.match(tauriAiCommandSource, /app\.emit\("ai-agent-event", &payload\)/);
|
||||
});
|
||||
|
||||
test("frontend ignores AI agent stream events from other sessions", () => {
|
||||
assert.match(tauriBackendSource, /type TauriAgentEvent = AgentEvent & \{\s*session_id\?: string;\s*\};/);
|
||||
assert.match(tauriBackendSource, /listen<TauriAgentEvent>\("ai-agent-event"/);
|
||||
assert.match(tauriBackendSource, /if \(payload\.session_id && payload\.session_id !== sessionId\) return;/);
|
||||
assert.match(tauriBackendSource, /onEvent\(payload\);/);
|
||||
});
|
||||
@@ -123,6 +123,13 @@ use dbx_core::agent_loop::{run_agent_loop, AgentLoopContext};
|
||||
use dbx_core::ai_cli_agent::CliAgentCommandSpec;
|
||||
use dbx_core::models::connection::DatabaseType;
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
struct AiAgentEventPayload {
|
||||
session_id: String,
|
||||
#[serde(flatten)]
|
||||
event: AgentEvent,
|
||||
}
|
||||
|
||||
#[tauri::command]
|
||||
pub async fn ai_cancel_stream(session_id: String) -> Result<bool, String> {
|
||||
Ok(dbx_core::ai::cancel_stream(&session_id).await)
|
||||
@@ -214,8 +221,10 @@ pub async fn ai_agent_stream(
|
||||
&agent_ctx,
|
||||
{
|
||||
let app = app.clone();
|
||||
let event_session_id = session_id.clone();
|
||||
move |event: AgentEvent| {
|
||||
let _ = app.emit("ai-agent-event", &event);
|
||||
let payload = AiAgentEventPayload { session_id: event_session_id.clone(), event };
|
||||
let _ = app.emit("ai-agent-event", &payload);
|
||||
}
|
||||
},
|
||||
&cancelled,
|
||||
|
||||
Reference in New Issue
Block a user