diff --git a/server/src/__tests__/helpers/test-db.ts b/server/src/__tests__/helpers/test-db.ts index dd9636dd..99da1d07 100644 --- a/server/src/__tests__/helpers/test-db.ts +++ b/server/src/__tests__/helpers/test-db.ts @@ -160,6 +160,18 @@ export const createTestDb = async () => { created_at TEXT NOT NULL ); + CREATE TABLE IF NOT EXISTS retrieval_runs ( + id TEXT PRIMARY KEY, + answer_run_id TEXT, + thread_id TEXT, + message_id TEXT, + turn_index INTEGER, + query TEXT NOT NULL, + hits TEXT NOT NULL, + duration_ms INTEGER, + created_at TEXT NOT NULL + ); + -- 週次レポート CREATE TABLE IF NOT EXISTS weekly_reports ( id TEXT PRIMARY KEY, diff --git a/server/src/db/migrations/0028_retrieval_runs.sql b/server/src/db/migrations/0028_retrieval_runs.sql new file mode 100644 index 00000000..1de36d98 --- /dev/null +++ b/server/src/db/migrations/0028_retrieval_runs.sql @@ -0,0 +1,15 @@ +CREATE TABLE retrieval_runs ( + id TEXT PRIMARY KEY, + answer_run_id TEXT, + thread_id TEXT, + message_id TEXT, + turn_index INTEGER, + query TEXT NOT NULL, + hits TEXT NOT NULL, + duration_ms INTEGER, + created_at TEXT NOT NULL +); + +CREATE INDEX idx_retrieval_runs_answer_run ON retrieval_runs(answer_run_id); +CREATE INDEX idx_retrieval_runs_thread ON retrieval_runs(thread_id, turn_index); +CREATE INDEX idx_retrieval_runs_created_at ON retrieval_runs(created_at); diff --git a/server/src/db/schema.ts b/server/src/db/schema.ts index b8585a3d..b86486d8 100644 --- a/server/src/db/schema.ts +++ b/server/src/db/schema.ts @@ -214,6 +214,21 @@ export const llmUsage = sqliteTable("llm_usage", { export type LlmUsage = typeof llmUsage.$inferSelect; export type NewLlmUsage = typeof llmUsage.$inferInsert; +export const retrievalRuns = sqliteTable("retrieval_runs", { + id: text("id").primaryKey(), + answerRunId: text("answer_run_id"), + threadId: text("thread_id"), + messageId: text("message_id"), + turnIndex: integer("turn_index"), + query: text("query").notNull(), + hits: text("hits").notNull(), // JSON(RetrievalHit[]) + durationMs: integer("duration_ms"), + createdAt: text("created_at").notNull(), +}); + +export type RetrievalRun = typeof retrievalRuns.$inferSelect; +export type NewRetrievalRun = typeof retrievalRuns.$inferInsert; + // 週次レポート(数値集計 + LLM ハイライト要約の恒久記録) export const weeklyReports = sqliteTable("weekly_reports", { id: text("id").primaryKey(), diff --git a/server/src/lib/wait-until.ts b/server/src/lib/wait-until.ts new file mode 100644 index 00000000..d48c264f --- /dev/null +++ b/server/src/lib/wait-until.ts @@ -0,0 +1,7 @@ +// cloudflare:workers は workerd 専用モジュール。top-level import すると +// このファイルを import 経由で読む Node 実行の eval スクリプトが落ちる +export const waitUntilInBackground = (promise: Promise) => { + void import("cloudflare:workers") + .then(({ waitUntil }) => waitUntil(promise)) + .catch(() => {}); +}; diff --git a/server/src/mastra/request-context.ts b/server/src/mastra/request-context.ts index 8d1fd032..19872cb5 100644 --- a/server/src/mastra/request-context.ts +++ b/server/src/mastra/request-context.ts @@ -47,6 +47,8 @@ export const createRequestContext = (values: MastraRequestContextType) => { if (values.usageTurnIndex !== undefined) { requestContext.set("usageTurnIndex", values.usageTurnIndex); } + requestContext.set("answerRunId", crypto.randomUUID()); + requestContext.set("retrievalTracePending", []); if (values.voiceFindings) { requestContext.set("voiceFindings", values.voiceFindings); } diff --git a/server/src/repository/retrieval-run-repository.test.ts b/server/src/repository/retrieval-run-repository.test.ts new file mode 100644 index 00000000..8729ee10 --- /dev/null +++ b/server/src/repository/retrieval-run-repository.test.ts @@ -0,0 +1,68 @@ +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { createTestDb, type TestDb } from "~/__tests__/helpers/test-db"; +import { retrievalRuns } from "~/db"; + +const { testDbHolder } = vi.hoisted(() => ({ + testDbHolder: { db: null as TestDb | null }, +})); + +vi.mock("~/db", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + createDb: () => testDbHolder.db, + }; +}); + +const { retrievalRunRepository } = await import("./retrieval-run-repository"); + +const d1 = {} as D1Database; + +const run = (id: string, createdAt: string) => ({ + id, + answerRunId: `ar-${id}`, + threadId: "t-1", + messageId: "m-1", + turnIndex: 1, + query: "村営バスの時刻", + hits: "[]", + durationMs: 100, + createdAt, +}); + +let db: TestDb; + +beforeEach(async () => { + db = await createTestDb(); + testDbHolder.db = db; +}); + +describe("deleteCreatedBefore", () => { + it("期限より前の検索記録だけを削除する", async () => { + await db + .insert(retrievalRuns) + .values([ + run("old", "2029-01-01T00:00:00Z"), + run("new", "2031-01-01T00:00:00Z"), + ]); + + const deleted = await retrievalRunRepository.deleteCreatedBefore( + d1, + "2030-01-01T00:00:00Z", + ); + + expect(deleted).toBe(1); + const remaining = await db.select().from(retrievalRuns).all(); + expect(remaining.map((r) => r.id)).toEqual(["new"]); + }); + + it("対象が無ければ 0 を返す", async () => { + const deleted = await retrievalRunRepository.deleteCreatedBefore( + d1, + "2030-01-01T00:00:00Z", + ); + + expect(deleted).toBe(0); + }); +}); diff --git a/server/src/repository/retrieval-run-repository.ts b/server/src/repository/retrieval-run-repository.ts new file mode 100644 index 00000000..0ca42f07 --- /dev/null +++ b/server/src/repository/retrieval-run-repository.ts @@ -0,0 +1,16 @@ +import { sql } from "drizzle-orm"; + +import { createDb, retrievalRuns } from "~/db"; +import { deleteWithCount } from "./delete-with-count"; + +export const retrievalRunRepository = { + async deleteCreatedBefore(d1: D1Database, cutoff: string) { + const db = createDb(d1); + + return deleteWithCount( + db, + retrievalRuns, + sql`datetime(${retrievalRuns.createdAt}) < datetime(${cutoff})`, + ); + }, +}; diff --git a/server/src/routes/threads/chat.test.ts b/server/src/routes/threads/chat.test.ts index e9459b24..20d5398f 100644 --- a/server/src/routes/threads/chat.test.ts +++ b/server/src/routes/threads/chat.test.ts @@ -53,6 +53,14 @@ vi.mock("~/mastra/request-context", () => ({ createRequestContext: vi.fn(() => ({})), })); +const { mockLinkRetrievalRunsToMessage } = vi.hoisted(() => ({ + mockLinkRetrievalRunsToMessage: vi.fn(), +})); + +vi.mock("~/services/knowledge/retrieval-trace", () => ({ + linkRetrievalRunsToMessage: mockLinkRetrievalRunsToMessage, +})); + vi.mock("@mastra/memory", () => ({ Memory: vi.fn(function () { return { @@ -269,6 +277,49 @@ describe("chatRoutes: POST /:threadId/chat", () => { ); }); + it("ストリーム完了時に assistant メッセージ ID を retrieval run へ紐付ける", async () => { + useAnonAuth(); + mockGetThreadById.mockResolvedValue(ownThread); + + await routes.request( + buildReq("/thread-1/chat", validBody), + undefined, + mockEnv, + ); + + const { onFinish } = mockHandleChatStream.mock.calls[0][0].params; + onFinish({ + totalUsage: { inputTokens: 1 }, + response: { + uiMessages: [ + { id: "user-1", role: "user" }, + { id: "assistant-1", role: "assistant" }, + ], + }, + }); + + expect(mockLinkRetrievalRunsToMessage).toHaveBeenCalledWith( + expect.anything(), + "assistant-1", + ); + }); + + it("assistant メッセージが無ければ紐付けしない", async () => { + useAnonAuth(); + mockGetThreadById.mockResolvedValue(ownThread); + + await routes.request( + buildReq("/thread-1/chat", validBody), + undefined, + mockEnv, + ); + + const { onFinish } = mockHandleChatStream.mock.calls[0][0].params; + onFinish({ totalUsage: { inputTokens: 1 } }); + + expect(mockLinkRetrievalRunsToMessage).not.toHaveBeenCalled(); + }); + it("widget- prefix の resourceId では usage を platform=widget で記録する", async () => { useAnonAuth(WIDGET_RES_ID); mockGetThreadById.mockResolvedValue({ diff --git a/server/src/routes/threads/chat.ts b/server/src/routes/threads/chat.ts index 6859ee0d..d4c0b07a 100644 --- a/server/src/routes/threads/chat.ts +++ b/server/src/routes/threads/chat.ts @@ -15,6 +15,7 @@ import type { ThreadVariables } from "~/middleware/require-thread-access"; import { requireThreadAccess } from "~/middleware/require-thread-access"; import { widgetSiteRepository } from "~/repository/widget-site-repository"; import { nextTurnIndex, recordLlmUsage } from "~/services/analytics/llm-usage"; +import { linkRetrievalRunsToMessage } from "~/services/knowledge/retrieval-trace"; export const chatRoutes = new OpenAPIHono<{ Bindings: CloudflareBindings; @@ -34,6 +35,16 @@ const ChatSendRequestSchema = z.object({ isGreeting: z.boolean().optional(), }); +const extractAssistantMessageId = (event: unknown) => { + const uiMessages = + ( + event as { + response?: { uiMessages?: Array<{ id?: string; role?: string }> }; + } + ).response?.uiMessages ?? []; + return uiMessages.findLast((message) => message.role === "assistant")?.id; +}; + const chatRoute = createRoute({ method: "post", path: "/{threadId}/chat", @@ -153,7 +164,7 @@ chatRoutes.openapi(chatRoute, async (c) => { thread: threadId, }, // onFinish はレスポンス返却後に発火するため waitUntil で記録を完了させる - onFinish: (event) => + onFinish: (event) => { waitUntil( recordLlmUsage(c.env.DB, { model: event.model?.modelId ?? primaryModelId(modelConfig), @@ -166,6 +177,13 @@ chatRoutes.openapi(chatRoute, async (c) => { turnIndex, durationMs: Date.now() - startedAt, }), - ), + ); + const assistantMessageId = extractAssistantMessageId(event); + if (assistantMessageId) { + waitUntil( + linkRetrievalRunsToMessage(requestContext, assistantMessageId), + ); + } + }, }); }); diff --git a/server/src/services/analytics/llm-usage.ts b/server/src/services/analytics/llm-usage.ts index e6841d92..adb1233e 100644 --- a/server/src/services/analytics/llm-usage.ts +++ b/server/src/services/analytics/llm-usage.ts @@ -2,6 +2,7 @@ import type { RequestContext } from "@mastra/core/request-context"; import type { MastraOnFinishCallbackArgs } from "@mastra/core/stream"; import { calcCostUsd, type LlmServiceTier } from "~/lib/llm-pricing"; import { logger } from "~/lib/logger"; +import { waitUntilInBackground } from "~/lib/wait-until"; import { llmUsageRepository } from "~/repository/llm-usage-repository"; export type LlmUsagePlatform = "web" | "line" | "lp" | "widget" | "voice"; @@ -183,11 +184,7 @@ export const usageRecordingOptions = durationMs: Date.now() - startedAt, ...contextAttributes(requestContext), }); - // cloudflare:workers は workerd 専用モジュール。top-level import すると - // このファイルを import 経由で読む Node 実行の eval スクリプトが落ちる - void import("cloudflare:workers") - .then(({ waitUntil }) => waitUntil(recording)) - .catch(() => {}); + waitUntilInBackground(recording); return recording; }, }; diff --git a/server/src/services/data-retention.test.ts b/server/src/services/data-retention.test.ts index 672bac77..9437cc7e 100644 --- a/server/src/services/data-retention.test.ts +++ b/server/src/services/data-retention.test.ts @@ -10,6 +10,7 @@ import { messageFeedback, pollSubmissions, polls, + retrievalRuns, threadPersonaStatus, } from "~/db"; @@ -252,6 +253,29 @@ describe("runDataRetention", () => { expect(remaining.map((r) => r.id)).toEqual(["fresh"]); }); + it("retrieval_runs のうち 90 日より前のものだけ削除する", async () => { + const db = testDbHolder.db as TestDb; + await db.insert(retrievalRuns).values([ + { + id: "old", + query: "q", + hits: "[]", + createdAt: daysAgo(91), + }, + { + id: "fresh", + query: "q", + hits: "[]", + createdAt: daysAgo(1), + }, + ]); + + await runDataRetention(env, { now: NOW }); + + const remaining = await db.select().from(retrievalRuns); + expect(remaining.map((r) => r.id)).toEqual(["fresh"]); + }); + it("data_retention_logs のうち 1095 日より前のものを自身も削除する", async () => { const db = testDbHolder.db as TestDb; await db.insert(dataRetentionLogs).values([ @@ -320,6 +344,7 @@ describe("runDataRetention", () => { "message_feedback", "llm_usage", "poll_submissions", + "retrieval_runs", "data_retention_logs", ]), ); @@ -337,6 +362,7 @@ describe("runDataRetention", () => { { table: "message_feedback", deletedCount: 0 }, { table: "llm_usage", deletedCount: 0 }, { table: "poll_submissions", deletedCount: 0 }, + { table: "retrieval_runs", deletedCount: 0 }, { table: "data_retention_logs", deletedCount: 0 }, ]); }); @@ -344,7 +370,7 @@ describe("runDataRetention", () => { it("options を省略すると現在時刻を基準に実行する", async () => { const results = await runDataRetention(env); - expect(results).toHaveLength(8); + expect(results).toHaveLength(9); expect(results.every((r) => r.deletedCount === 0)).toBe(true); }); diff --git a/server/src/services/data-retention.ts b/server/src/services/data-retention.ts index 13fc65ef..6df29673 100644 --- a/server/src/services/data-retention.ts +++ b/server/src/services/data-retention.ts @@ -7,6 +7,7 @@ import { mastraMessageRepository } from "~/repository/mastra-message-repository" import { mastraResourceRepository } from "~/repository/mastra-resource-repository"; import { mastraThreadRepository } from "~/repository/mastra-thread-repository"; import { pollRepository } from "~/repository/poll-repository"; +import { retrievalRunRepository } from "~/repository/retrieval-run-repository"; import { threadPersonaStatusRepository } from "~/repository/thread-persona-status-repository"; const RETENTION_DAYS = { @@ -16,6 +17,7 @@ const RETENTION_DAYS = { message_feedback: 180, llm_usage: 180, // 週次レポートに集計が恒久保存されるため raw は短期で良い poll_submissions: 365, + retrieval_runs: 90, data_retention_logs: 1095, } as const; @@ -103,6 +105,14 @@ export const runDataRetention = async ( ), }); + results.push({ + table: "retrieval_runs", + deletedCount: await retrievalRunRepository.deleteCreatedBefore( + env.DB, + cutoff(now, RETENTION_DAYS.retrieval_runs), + ), + }); + results.push({ table: "data_retention_logs", deletedCount: await dataRetentionLogRepository.deleteExecutedBefore( diff --git a/server/src/services/knowledge/embedding.test.ts b/server/src/services/knowledge/embedding.test.ts index 2fb2fa29..9144eaef 100644 --- a/server/src/services/knowledge/embedding.test.ts +++ b/server/src/services/knowledge/embedding.test.ts @@ -131,6 +131,25 @@ describe("processKnowledgeFile", () => { expect(upsertArg[1].metadata.subsection).toBe("Sub"); }); + it("chunk 本文の SHA-256 を contentHash として付与する", async () => { + const bodyA = "a".repeat(100); + setChunks([bodyA, bodyA, "b".repeat(100)], [{}, {}, {}]); + vi.mocked(embedMany).mockResolvedValueOnce({ + embeddings: [fakeEmbedding(), fakeEmbedding(), fakeEmbedding()], + } as never); + + const vectorize = buildVectorize(); + await processKnowledgeFile("guide.md", "x", vectorize, "key"); + + const upsertArg = vi.mocked(vectorize.upsert).mock.calls[0]?.[0] as Array<{ + metadata: Record; + }>; + const hashes = upsertArg.map((v) => v.metadata.contentHash); + expect(hashes[0]).toMatch(/^[0-9a-f]{64}$/); + expect(hashes[0]).toBe(hashes[1]); + expect(hashes[2]).not.toBe(hashes[0]); + }); + it("strip されたヘッダー情報を embedding テキストの先頭にプレフィックスとして復元する", async () => { const bodyA = "a".repeat(100); const bodyB = "b".repeat(100); diff --git a/server/src/services/knowledge/embedding.ts b/server/src/services/knowledge/embedding.ts index 87e2cd31..616cfbee 100644 --- a/server/src/services/knowledge/embedding.ts +++ b/server/src/services/knowledge/embedding.ts @@ -21,9 +21,20 @@ type ChunkMetadata = { section?: string; subsection?: string; content: string; + contentHash: string; [key: string]: string | number | boolean | string[] | undefined; }; +const hashText = async (text: string) => { + const digest = await crypto.subtle.digest( + "SHA-256", + new TextEncoder().encode(text), + ); + return [...new Uint8Array(digest)] + .map((byte) => byte.toString(16).padStart(2, "0")) + .join(""); +}; + type VectorData = { id: string; values: number[]; @@ -85,19 +96,22 @@ const chunkDocument = async ( .join(" > "); return prefix ? `${prefix}\n\n${allTexts[i]}` : allTexts[i]; }); - const metadata: ChunkMetadata[] = filteredIndices.map((i, idx) => { - const chunkMeta = chunkMetadataList[i] as - | Record - | undefined; - return { - ...frontmatter, - source: filename, - title: chunkMeta?.title as string | undefined, - section: chunkMeta?.section as string | undefined, - subsection: chunkMeta?.subsection as string | undefined, - content: texts[idx], - }; - }); + const metadata: ChunkMetadata[] = await Promise.all( + filteredIndices.map(async (i, idx) => { + const chunkMeta = chunkMetadataList[i] as + | Record + | undefined; + return { + ...frontmatter, + source: filename, + title: chunkMeta?.title as string | undefined, + section: chunkMeta?.section as string | undefined, + subsection: chunkMeta?.subsection as string | undefined, + content: texts[idx], + contentHash: await hashText(texts[idx]), + }; + }), + ); logger.info( `[Knowledge Sync] ${filename}: ${allTexts.length} chunks -> ${texts.length} after filtering (min ${MIN_CHUNK_LENGTH} chars)`, diff --git a/server/src/services/knowledge/retrieval-trace.test.ts b/server/src/services/knowledge/retrieval-trace.test.ts new file mode 100644 index 00000000..83015138 --- /dev/null +++ b/server/src/services/knowledge/retrieval-trace.test.ts @@ -0,0 +1,211 @@ +import { RequestContext } from "@mastra/core/request-context"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { createTestDb, type TestDb } from "~/__tests__/helpers/test-db"; +import { retrievalRuns } from "~/db"; + +const { testDbHolder } = vi.hoisted(() => ({ + testDbHolder: { db: null as TestDb | null }, +})); + +const { loggerMock } = vi.hoisted(() => ({ + loggerMock: { info: vi.fn(), error: vi.fn(), warn: vi.fn(), debug: vi.fn() }, +})); + +vi.mock("~/db", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + createDb: () => testDbHolder.db, + }; +}); + +vi.mock("~/lib/logger", () => ({ + logger: loggerMock, +})); + +const { + recordRetrievalRun, + recordRetrievalRunInBackground, + linkRetrievalRunsToMessage, +} = await import("./retrieval-trace"); + +const d1 = {} as D1Database; + +const buildContext = (values: Record = {}) => { + const context = new RequestContext(); + context.set("db", d1); + for (const [key, value] of Object.entries(values)) { + context.set(key, value); + } + return context; +}; + +describe("recordRetrievalRun", () => { + let db: TestDb; + + beforeEach(async () => { + db = await createTestDb(); + testDbHolder.db = db; + vi.clearAllMocks(); + }); + + it("検索結果を requestContext の属性とともに保存する", async () => { + const context = buildContext({ + answerRunId: "run-1", + usageThreadId: "thread-1", + usageTurnIndex: 3, + }); + + await recordRetrievalRun(context, { + query: "村営バスの時刻", + durationMs: 120, + hits: [ + { + source: "bus/index.md", + title: "村営バス", + section: "時刻表", + score: 0.82, + rerankScore: 0.9, + contentHash: "abc123", + }, + ], + }); + + const rows = await db.select().from(retrievalRuns).all(); + expect(rows).toHaveLength(1); + expect(rows[0]).toMatchObject({ + answerRunId: "run-1", + threadId: "thread-1", + messageId: null, + turnIndex: 3, + query: "村営バスの時刻", + durationMs: 120, + }); + expect(JSON.parse(rows[0].hits)).toEqual([ + { + source: "bus/index.md", + title: "村営バス", + section: "時刻表", + score: 0.82, + rerankScore: 0.9, + contentHash: "abc123", + }, + ]); + }); + + it("0 件の検索も hits 空配列で保存する", async () => { + await recordRetrievalRun(buildContext({ answerRunId: "run-1" }), { + query: "存在しない情報", + hits: [], + }); + + const rows = await db.select().from(retrievalRuns).all(); + expect(rows).toHaveLength(1); + expect(JSON.parse(rows[0].hits)).toEqual([]); + }); + + it("requestContext に db が無ければ何もしない", async () => { + await recordRetrievalRun(new RequestContext(), { + query: "foo", + hits: [], + }); + + const rows = await db.select().from(retrievalRuns).all(); + expect(rows).toHaveLength(0); + }); + + it("保存に失敗しても throw せず警告ログを残す", async () => { + testDbHolder.db = null; + + await expect( + recordRetrievalRun(buildContext(), { query: "foo", hits: [] }), + ).resolves.toBeUndefined(); + expect(loggerMock.warn).toHaveBeenCalled(); + }); +}); + +describe("linkRetrievalRunsToMessage", () => { + let db: TestDb; + + beforeEach(async () => { + db = await createTestDb(); + testDbHolder.db = db; + vi.clearAllMocks(); + }); + + const insertRun = (values: { + id: string; + answerRunId: string; + messageId?: string; + }) => + db.insert(retrievalRuns).values({ + id: values.id, + answerRunId: values.answerRunId, + messageId: values.messageId, + query: "q", + hits: "[]", + createdAt: new Date().toISOString(), + }); + + it("同じ answer run の全検索行に message_id を後埋めする", async () => { + await insertRun({ id: "r1", answerRunId: "run-1" }); + await insertRun({ id: "r2", answerRunId: "run-1" }); + await insertRun({ id: "r3", answerRunId: "run-2" }); + + await linkRetrievalRunsToMessage( + buildContext({ answerRunId: "run-1" }), + "msg-1", + ); + + const rows = await db.select().from(retrievalRuns).all(); + const byId = Object.fromEntries(rows.map((r) => [r.id, r.messageId])); + expect(byId).toEqual({ r1: "msg-1", r2: "msg-1", r3: null }); + }); + + it("既に message_id が入っている行は上書きしない", async () => { + await insertRun({ id: "r1", answerRunId: "run-1", messageId: "msg-old" }); + + await linkRetrievalRunsToMessage( + buildContext({ answerRunId: "run-1" }), + "msg-new", + ); + + const rows = await db.select().from(retrievalRuns).all(); + expect(rows[0].messageId).toBe("msg-old"); + }); + + it("未完了の保存があれば完了を待ってから紐付ける", async () => { + const context = buildContext({ + answerRunId: "run-1", + retrievalTracePending: [], + }); + recordRetrievalRunInBackground(context, { query: "q", hits: [] }); + + await linkRetrievalRunsToMessage(context, "msg-1"); + + const rows = await db.select().from(retrievalRuns).all(); + expect(rows).toHaveLength(1); + expect(rows[0].messageId).toBe("msg-1"); + }); + + it("answerRunId が無ければ何もしない", async () => { + await insertRun({ id: "r1", answerRunId: "run-1" }); + + await linkRetrievalRunsToMessage(buildContext(), "msg-1"); + + const rows = await db.select().from(retrievalRuns).all(); + expect(rows[0].messageId).toBeNull(); + }); + + it("更新に失敗しても throw せず警告ログを残す", async () => { + testDbHolder.db = null; + + await expect( + linkRetrievalRunsToMessage( + buildContext({ answerRunId: "run-1" }), + "msg-1", + ), + ).resolves.toBeUndefined(); + expect(loggerMock.warn).toHaveBeenCalled(); + }); +}); diff --git a/server/src/services/knowledge/retrieval-trace.ts b/server/src/services/knowledge/retrieval-trace.ts new file mode 100644 index 00000000..0e37fa8b --- /dev/null +++ b/server/src/services/knowledge/retrieval-trace.ts @@ -0,0 +1,88 @@ +import type { RequestContext } from "@mastra/core/request-context"; +import { and, eq, isNull } from "drizzle-orm"; +import { createDb, retrievalRuns } from "~/db"; +import { logger } from "~/lib/logger"; +import { waitUntilInBackground } from "~/lib/wait-until"; + +export type RetrievalHit = { + source: string; + title?: string; + section?: string; + score: number; + rerankScore?: number; + contentHash?: string; +}; + +type RecordRetrievalRunParams = { + query: string; + hits: RetrievalHit[]; + durationMs?: number; +}; + +export const recordRetrievalRun = async ( + requestContext: RequestContext | undefined, + params: RecordRetrievalRunParams, +) => { + const d1 = requestContext?.get("db") as D1Database | undefined; + if (!d1) return; + try { + const db = createDb(d1); + await db.insert(retrievalRuns).values({ + id: crypto.randomUUID(), + answerRunId: requestContext?.get("answerRunId") as string | undefined, + threadId: requestContext?.get("usageThreadId") as string | undefined, + turnIndex: requestContext?.get("usageTurnIndex") as number | undefined, + query: params.query, + hits: JSON.stringify(params.hits), + durationMs: params.durationMs, + createdAt: new Date().toISOString(), + }); + } catch (error) { + logger.warn("[RetrievalTrace] failed to record run", { + error: String(error), + }); + } +}; + +export const recordRetrievalRunInBackground = ( + requestContext: RequestContext | undefined, + params: RecordRetrievalRunParams, +) => { + const recording = recordRetrievalRun(requestContext, params); + const pending = requestContext?.get("retrievalTracePending") as + | Promise[] + | undefined; + pending?.push(recording); + waitUntilInBackground(recording); +}; + +export const linkRetrievalRunsToMessage = async ( + requestContext: RequestContext | undefined, + messageId: string, +) => { + const d1 = requestContext?.get("db") as D1Database | undefined; + const answerRunId = requestContext?.get("answerRunId") as string | undefined; + if (!d1 || !answerRunId) return; + const pending = requestContext?.get("retrievalTracePending") as + | Promise[] + | undefined; + if (pending?.length) { + await Promise.allSettled(pending); + } + try { + const db = createDb(d1); + await db + .update(retrievalRuns) + .set({ messageId }) + .where( + and( + eq(retrievalRuns.answerRunId, answerRunId), + isNull(retrievalRuns.messageId), + ), + ); + } catch (error) { + logger.warn("[RetrievalTrace] failed to link message", { + error: String(error), + }); + } +}; diff --git a/server/src/services/knowledge/search.test.ts b/server/src/services/knowledge/search.test.ts index 953c1688..1327303e 100644 --- a/server/src/services/knowledge/search.test.ts +++ b/server/src/services/knowledge/search.test.ts @@ -32,9 +32,16 @@ vi.mock("~/lib/logger", () => ({ logger: { info: vi.fn(), warn: vi.fn(), error: vi.fn() }, })); +vi.mock("~/services/knowledge/retrieval-trace", () => ({ + recordRetrievalRunInBackground: vi.fn(), +})); + const { embed } = await import("ai"); const { rerankWithScorer } = await import("@mastra/rag"); const { logger } = await import("~/lib/logger"); +const { recordRetrievalRunInBackground } = await import( + "~/services/knowledge/retrieval-trace" +); const { searchKnowledge } = await import("./search"); const buildVectorize = () => @@ -48,6 +55,7 @@ beforeEach(() => { vi.mocked(embed).mockReset(); vi.mocked(rerankWithScorer).mockReset(); vi.mocked(logger.error).mockReset(); + vi.mocked(recordRetrievalRunInBackground).mockReset(); }); describe("searchKnowledge", () => { @@ -142,6 +150,63 @@ describe("searchKnowledge", () => { expect(rerankArg.results[0].metadata.source).toBe("doc.md"); }); + it("rerank 後の hits をトレースに記録する", async () => { + vi.mocked(embed).mockResolvedValueOnce({ embedding: [0.1] } as never); + const vectorize = buildVectorize(); + vi.mocked(vectorize.query).mockResolvedValueOnce({ + matches: [{ id: "v1", score: 0.8, metadata: { content: "本文1" } }], + } as never); + vi.mocked(rerankWithScorer).mockResolvedValueOnce([ + { + score: 0.95, + result: { + id: "v1", + score: 0.8, + metadata: { + content: "本文1", + source: "doc.md", + title: "Tタイトル", + section: "Sセク", + contentHash: "hash1", + }, + }, + }, + ] as never); + + await searchKnowledge("クエリ", vectorize, "key"); + + expect(recordRetrievalRunInBackground).toHaveBeenCalledWith(undefined, { + query: "クエリ", + hits: [ + { + source: "doc.md", + title: "Tタイトル", + section: "Sセク", + score: 0.8, + rerankScore: 0.95, + contentHash: "hash1", + }, + ], + durationMs: expect.any(Number), + }); + }); + + it("0 件の検索も hits 空配列でトレースに記録する", async () => { + vi.mocked(embed).mockResolvedValueOnce({ embedding: [0.1] } as never); + const vectorize = buildVectorize(); + vi.mocked(vectorize.query).mockResolvedValueOnce({ + matches: [], + } as never); + + await searchKnowledge("見つからないクエリ", vectorize, "key"); + + expect(recordRetrievalRunInBackground).toHaveBeenCalledWith(undefined, { + query: "見つからないクエリ", + hits: [], + durationMs: expect.any(Number), + }); + }); + it("metadata 欠落フィールドは unknown / 空文字で埋める", async () => { vi.mocked(embed).mockResolvedValueOnce({ embedding: [0.1] } as never); const vectorize = buildVectorize(); diff --git a/server/src/services/knowledge/search.ts b/server/src/services/knowledge/search.ts index 0f642281..efcb4d6b 100644 --- a/server/src/services/knowledge/search.ts +++ b/server/src/services/knowledge/search.ts @@ -14,6 +14,7 @@ import { recordUsageFromContext, withUsageRecording, } from "~/services/analytics/llm-usage"; +import { recordRetrievalRunInBackground } from "~/services/knowledge/retrieval-trace"; const EMBEDDING_DIMENSIONS = 1536; @@ -77,6 +78,7 @@ export const searchKnowledge = async ( apiKey: string, requestContext?: RequestContext, ): Promise => { + const startedAt = Date.now(); try { logger.info("[Knowledge] search", { query }); const google = createGoogleGenerativeAI({ apiKey }); @@ -105,6 +107,11 @@ export const searchKnowledge = async ( }); if (!results.matches || results.matches.length === 0) { + recordRetrievalRunInBackground(requestContext, { + query, + hits: [], + durationMs: Date.now() - startedAt, + }); return { results: [], }; @@ -125,6 +132,7 @@ export const searchKnowledge = async ( url: metadata?.url as string | undefined, date: metadata?.date as string | undefined, dateType: metadata?.date_type as string | undefined, + contentHash: metadata?.contentHash as string | undefined, }, }; }); @@ -155,6 +163,21 @@ export const searchKnowledge = async ( dateType: r.result.metadata?.dateType as string | undefined, })); + recordRetrievalRunInBackground(requestContext, { + query, + hits: knowledgeResults.map((result, i) => ({ + source: result.source, + title: result.title, + section: result.section, + score: rerankedResults[i].result.score, + rerankScore: result.score, + contentHash: rerankedResults[i].result.metadata?.contentHash as + | string + | undefined, + })), + durationMs: Date.now() - startedAt, + }); + logger.info("[Knowledge] search result", { query, hits: knowledgeResults