diff --git a/e2e/citation-dedupe.e2e.ts b/e2e/citation-dedupe.e2e.ts index fa34f4a..62b6f76 100644 --- a/e2e/citation-dedupe.e2e.ts +++ b/e2e/citation-dedupe.e2e.ts @@ -276,3 +276,103 @@ test("draws citation regions instead of a full-page highlight", async ({ ) } }) + +test("moves a deduplicated cited page to the top from the same source", async ({ + context, + page, +}) => { + await context.addCookies([ + { + name: "better-auth.session_token", + value: "playwright", + url: "http://localhost:3000", + }, + ]) + await page.setViewportSize({ width: 1280, height: 832 }) + await page.route("**/api/sources/source_spacex/chunks**", async (route) => { + const pageAsset = { + pageNumber: 26, + assetUrl: "/images/knowhere/logo-icon.png", + contentType: "image/png", + width: 1000, + height: 1400, + } + await route.fulfill({ + status: 200, + contentType: "application/json", + body: JSON.stringify({ + chunks: [ + { + chunkId: "page_1", + documentId: "doc_spacex", + sectionPath: "Page 1", + type: "page", + content: "Content from page 1.", + readableContent: "Content from page 1.", + pageNums: [1], + pageAssets: [{ ...pageAsset, pageNumber: 1 }], + sourceTitle: "spacex-s1.pdf", + }, + { + chunkId: "page_26_first_section", + documentId: "doc_spacex", + sectionPath: "Page 26 / First section", + type: "page", + content: "Other content from page 26.", + readableContent: "Other content from page 26.", + pageNums: [26], + pageAssets: [pageAsset], + sourceTitle: "spacex-s1.pdf", + }, + { + chunkId: "page_26_cited_section", + documentId: "doc_spacex", + sectionPath: "Page 26", + type: "page", + content: "Revenue evidence on page 26.", + readableContent: "Revenue evidence on page 26.", + pageNums: [26], + pageAssets: [pageAsset], + sourceTitle: "spacex-s1.pdf", + }, + ], + pagination: { + page: 1, + pageSize: 50, + total: 3, + totalPages: 1, + }, + }), + }) + }) + + await page.goto("/e2e/citation-same-page") + await page + .getByTestId("desktop-sources-panel") + .getByRole("button", { name: "Open spacex-s1.pdf parsed chunks" }) + .click() + const chunksPanel = page.getByTestId("desktop-chunks-panel") + await chunksPanel.getByRole("button", { name: "List" }).click() + await expect + .poll(() => + chunksPanel + .locator('[data-index="0"]') + .getAttribute("data-chunk-id"), + ) + .toBe("page_1") + + await page + .getByTestId("desktop-chat-panel") + .getByTestId("citation-chip") + .first() + .click() + + await expect + .poll(() => + chunksPanel + .locator('[data-index="0"]') + .getAttribute("data-chunk-id"), + ) + .toBe("page_26_first_section") + await expect(page.getByTestId("citation-region-highlight")).toHaveCount(1) +}) diff --git a/src/app/e2e/citation-same-page/page.tsx b/src/app/e2e/citation-same-page/page.tsx index d2c16f9..1080aa7 100644 --- a/src/app/e2e/citation-same-page/page.tsx +++ b/src/app/e2e/citation-same-page/page.tsx @@ -34,6 +34,7 @@ const chatMessages: ChatMessageView[] = [ { chunkType: "page", score: 0.91, + content: "Revenue evidence on page 26.", pageCitationPageNumber: 26, highlightRegions: [{ x: 0.12, y: 0.18, w: 0.46, h: 0.08 }], source: { diff --git a/src/components/chunks-panel-state.test.ts b/src/components/chunks-panel-state.test.ts index febfd71..447dff2 100644 --- a/src/components/chunks-panel-state.test.ts +++ b/src/components/chunks-panel-state.test.ts @@ -208,6 +208,44 @@ describe("chunksPanelState", () => { ).toEqual(["page_4_first", "page_5"]) }) + it("focuses the retained page card when the cited logical chunk was deduplicated", () => { + const chunks: ParsedChunkView[] = [ + { + chunkId: "page_1", + type: "page", + content: "Page 1.", + sourceTitle: "cdn.pdf", + pageNums: [1], + }, + { + chunkId: "page_4_first_section", + type: "page", + content: "First section on page 4.", + sourceTitle: "cdn.pdf", + pageNums: [4], + }, + { + chunkId: "page_4_cited_section", + type: "page", + content: "Cited section on page 4.", + sourceTitle: "cdn.pdf", + pageNums: [4], + }, + ] + const displayedChunks = + chunksPanelState.getPageAssetChunksWithoutDuplicatePages(chunks) + + expect( + chunksPanelState + .getChunksWithFocusedFirst( + displayedChunks, + "page_4_cited_section", + 4, + ) + .map((chunk) => chunk.chunkId), + ).toEqual(["page_4_first_section", "page_1"]) + }) + it("keeps overlapping page-memory section chunks that share a boundary page", () => { const chunks: ParsedChunkView[] = [ { diff --git a/src/components/chunks-panel-state.ts b/src/components/chunks-panel-state.ts index 3afccbf..d688d08 100644 --- a/src/components/chunks-panel-state.ts +++ b/src/components/chunks-panel-state.ts @@ -63,8 +63,12 @@ function getChunksWithFocusedFirst( const orderedChunks = getChunksOrderedByPageNumber( dedupeChunksById(chunks), ) - const targetChunkId = - focusedChunkId ?? findChunkIdForPageNumber(orderedChunks, focusedPageNumber) + const hasVisibleFocusedChunk = + focusedChunkId !== null && + orderedChunks.some((chunk) => chunk.chunkId === focusedChunkId) + const targetChunkId = hasVisibleFocusedChunk + ? focusedChunkId + : findChunkIdForPageNumber(orderedChunks, focusedPageNumber) if (!targetChunkId) return orderedChunks const focusedIndex = orderedChunks.findIndex( diff --git a/src/domains/chat/index.test.ts b/src/domains/chat/index.test.ts index 003e436..1bc9b60 100644 --- a/src/domains/chat/index.test.ts +++ b/src/domains/chat/index.test.ts @@ -4,6 +4,8 @@ import type { KnowledgeGrepResponse, KnowledgeOutline, KnowledgeReadResponse, + RetrievalQueryParams, + RetrievalQueryResponse, RetrievalResult, } from "@ontos-ai/knowhere-sdk" import { Effect } from "effect" @@ -362,6 +364,155 @@ describe("answerQuestionWithRetrieval", () => { }); }); + it("queries at most four namespaces concurrently", async () => { + const namespaces = [ + "default", + "notebook-1", + "notebook-2", + "notebook-3", + "notebook-4", + "notebook-5", + ]; + const releases: Array<() => void> = []; + let activeCount = 0; + let maxActiveCount = 0; + const retrieval = { + query: vi.fn(async (params: RetrievalQueryParams) => { + const namespace = requireRetrievalQueryText( + params.namespace, + "namespace", + ); + const query = requireRetrievalQueryText(params.query, "query"); + activeCount += 1; + maxActiveCount = Math.max(maxActiveCount, activeCount); + await new Promise((resolve) => { + releases.push(() => { + activeCount -= 1; + resolve(); + }); + }); + return makeRetrievalQueryResponse(namespace, query); + }), + }; + const generateAnswer = vi.fn(async ({ searchSources }) => { + await searchSources({ query: "parallel evidence" }); + return makeHarnessRunResult("The answer is grounded."); + }); + + const answerPromise = Effect.runPromise( + answerQuestionWithRetrieval({ + question: "What do the documents say?", + namespace: "notebook-1", + namespaces, + sources: [makeSource()], + excludedSourceIds: [], + retrieval, + generateAnswer, + messages: [], + }), + ); + + await vi.waitFor(() => expect(retrieval.query).toHaveBeenCalledTimes(4)); + expect(maxActiveCount).toBe(4); + releases.splice(0).forEach((release) => release()); + + await vi.waitFor(() => expect(retrieval.query).toHaveBeenCalledTimes(6)); + expect(maxActiveCount).toBe(4); + releases.splice(0).forEach((release) => release()); + + await answerPromise; + expect(maxActiveCount).toBe(4); + }); + + it("retries rate-limited namespace queries sequentially", async () => { + const namespaces = ["default", "notebook-1", "notebook-2"]; + const attempts = new Map(); + const fallbackReleases: Array<() => void> = []; + let activeFallbackCount = 0; + let maxActiveFallbackCount = 0; + const retrieval = { + query: vi.fn(async (params: RetrievalQueryParams) => { + const namespace = requireRetrievalQueryText( + params.namespace, + "namespace", + ); + const query = requireRetrievalQueryText(params.query, "query"); + const attempt = (attempts.get(namespace) ?? 0) + 1; + attempts.set(namespace, attempt); + + if (namespace !== "default" && attempt === 1) { + throw Object.assign( + new Error( + "Too many concurrent requests. Please retry after 30 seconds.", + ), + { + statusCode: 429, + code: "RESOURCE_EXHAUSTED", + }, + ); + } + + if (namespace !== "default") { + activeFallbackCount += 1; + maxActiveFallbackCount = Math.max( + maxActiveFallbackCount, + activeFallbackCount, + ); + await new Promise((resolve) => { + fallbackReleases.push(() => { + activeFallbackCount -= 1; + resolve(); + }); + }); + } + + return makeRetrievalQueryResponse(namespace, query); + }), + }; + const generateAnswer = vi.fn(async ({ searchSources }) => { + await searchSources({ query: "rate-limited evidence" }); + return makeHarnessRunResult("The answer is grounded."); + }); + + const answerPromise = Effect.runPromise( + answerQuestionWithRetrieval({ + question: "What do the documents say?", + namespace: "notebook-1", + namespaces, + sources: [makeSource()], + excludedSourceIds: [], + retrieval, + generateAnswer, + messages: [], + }), + ); + + await vi.waitFor(() => expect(retrieval.query).toHaveBeenCalledTimes(4)); + expect(maxActiveFallbackCount).toBe(1); + expect(attempts.get("notebook-2")).toBe(1); + fallbackReleases.shift()?.(); + + await vi.waitFor(() => expect(retrieval.query).toHaveBeenCalledTimes(5)); + expect(maxActiveFallbackCount).toBe(1); + fallbackReleases.shift()?.(); + + await answerPromise; + expect(attempts).toEqual( + new Map([ + ["default", 1], + ["notebook-1", 2], + ["notebook-2", 2], + ]), + ); + expect(loggerMock.warn).toHaveBeenCalledWith( + "chat-agent: searchSources rate limited; retrying failed namespaces sequentially", + { + namespaceCount: 2, + namespaces: ["notebook-1", "notebook-2"], + }, + ); + }); + it("bounds merged retrieval evidence before passing it to the answer agent", async () => { const defaultResults = Array.from({ length: 40 }, (_, index) => makeRetrievalResult({ @@ -3115,6 +3266,33 @@ function makeRetrievalResult( }; } +function makeRetrievalQueryResponse( + namespace: string, + query: string, +): RetrievalQueryResponse { + return { + results: [makeRetrievalResult()], + evidenceText: `Evidence from ${namespace}`, + referencedChunks: [], + namespace, + query, + routerUsed: "workflow_single_step", + answerText: null, + stopReason: "answer_done", + failureReason: null, + }; +} + +function requireRetrievalQueryText( + value: string | undefined, + name: string, +): string { + if (typeof value !== "string" || value.length === 0) { + throw new Error(`Expected retrieval ${name}.`); + } + return value; +} + function makeSource(overrides: Partial = {}): Source { return { id: "source_1", diff --git a/src/domains/chat/index.ts b/src/domains/chat/index.ts index 446b591..d2b7208 100644 --- a/src/domains/chat/index.ts +++ b/src/domains/chat/index.ts @@ -48,6 +48,7 @@ import { notebookKnowhereTools } from "./knowhere-tools" const DEFAULT_TOP_K = 8 const MAX_AGENTIC_TOP_K = 12 +const MAX_CONCURRENT_RETRIEVAL_NAMESPACES = 4 const MAX_AGENTIC_MERGED_RESULT_COUNT = 24 const MAX_AGENTIC_MERGED_REFERENCED_CHUNK_COUNT = 24 const MAX_AGENTIC_MERGED_TEXT_CHARS = 12_000 @@ -102,6 +103,20 @@ type AgenticMergedEvidenceLimits = { readonly referencedChunkCount: number } +type RetrievalNamespaceQueryMode = "concurrent" | "rate_limit_fallback" + +type RetrievalNamespaceQueryOutcome = + | { + readonly status: "success" + readonly namespace: string + readonly response: RetrievalQueryResponse + } + | { + readonly status: "failure" + readonly namespace: string + readonly error: unknown + } + export type { AnswerQuestionInput, AnswerQuestionResult, @@ -143,58 +158,70 @@ export const answerQuestionWithRetrieval = ( const queryResponses: RetrievalQueryResponse[] = [] const queryFailures: unknown[] = [] - for (const namespace of namespaces) { - const retrievalQueryParams = buildRetrievalQueryParams({ - input: queryInput, + const concurrentOutcomes = await queryRetrievalNamespaces({ + namespaces, + queryInput, + answerInput: input, + fallbackQuestion: question, + retrievalPlan, + concurrency: MAX_CONCURRENT_RETRIEVAL_NAMESPACES, + mode: "concurrent", + }) + const rateLimitedOutcomes = concurrentOutcomes.filter( + (outcome) => + outcome.status === "failure" && isRateLimitError(outcome.error), + ) + let finalOutcomes = concurrentOutcomes + + if (rateLimitedOutcomes.length > 0) { + logger.warn( + "chat-agent: searchSources rate limited; retrying failed namespaces sequentially", + { + namespaceCount: rateLimitedOutcomes.length, + namespaces: rateLimitedOutcomes.map((outcome) => outcome.namespace), + }, + ) + const fallbackOutcomes = await queryRetrievalNamespaces({ + namespaces: rateLimitedOutcomes.map((outcome) => outcome.namespace), + queryInput, + answerInput: input, fallbackQuestion: question, - namespace, - useAgentic: input.useAgentic ?? true, - sources: input.sources, - excludedSourceIds: input.excludedSourceIds, - }) - logger.info("chat-agent: searchSources start", { - namespace, - query: retrievalQueryParams.query, - topK: retrievalQueryParams.topK, - useAgentic: retrievalQueryParams.useAgentic, - dataType: retrievalQueryParams.dataType ?? null, - signalPathCount: retrievalQueryParams.signalPaths?.length ?? 0, - filterMode: retrievalQueryParams.filterMode ?? null, - threshold: retrievalQueryParams.threshold ?? null, - targetContent: retrievalPlan.targetContent, - purpose: retrievalPlan.purpose, + retrievalPlan, + concurrency: 1, + mode: "rate_limit_fallback", }) + const fallbackByNamespace = new Map( + fallbackOutcomes.map( + (outcome): readonly [string, RetrievalNamespaceQueryOutcome] => [ + outcome.namespace, + outcome, + ], + ), + ) + finalOutcomes = concurrentOutcomes.map((outcome) => + outcome.status === "failure" && isRateLimitError(outcome.error) + ? (fallbackByNamespace.get(outcome.namespace) ?? outcome) + : outcome, + ) + } - try { - const response = await input.retrieval.query(retrievalQueryParams) - retrievalResponses.push(response) - queryResponses.push(response) - logger.info("chat-agent: searchSources ok", { - namespace, - query: response.query, - durationMs: Date.now() - startedAt, - resultCount: response.results.length, - referencedChunkCount: response.referencedChunks.length, - stopReason: response.stopReason ?? null, - failureReason: response.failureReason ?? null, - targetContent: retrievalPlan.targetContent, - }) - logger.info("chat-agent: knowhere query response", { - durationMs: Date.now() - startedAt, - response: formatKnowhereQueryResponseForLog(response), - }) - } catch (error) { - queryFailures.push(error) - logger.error("chat-agent: searchSources failed", { - namespace, - query: retrievalQueryParams.query, - durationMs: Date.now() - startedAt, - error: formatUnknownError(error), - targetContent: retrievalPlan.targetContent, - }) + for (const outcome of finalOutcomes) { + if (outcome.status === "success") { + retrievalResponses.push(outcome.response) + queryResponses.push(outcome.response) + } else { + queryFailures.push(outcome.error) } } + logger.info("chat-agent: searchSources batch complete", { + durationMs: Date.now() - startedAt, + namespaceCount: namespaces.length, + successCount: queryResponses.length, + failureCount: queryFailures.length, + rateLimitFallbackCount: rateLimitedOutcomes.length, + }) + if (queryResponses.length === 0) throw queryFailures[0] if ( queryFailures.length > 0 && @@ -674,6 +701,139 @@ function formatUnknownError(error: unknown): string { return String(error) } +function queryRetrievalNamespaces(input: { + readonly namespaces: readonly string[] + readonly queryInput: AgenticRetrievalQuery + readonly answerInput: AnswerQuestionInput + readonly fallbackQuestion: string + readonly retrievalPlan: AgenticRetrievalPlan + readonly concurrency: number + readonly mode: RetrievalNamespaceQueryMode +}): Promise { + return Effect.runPromise( + Effect.all( + input.namespaces.map((namespace) => + Effect.promise(() => + queryRetrievalNamespace({ + namespace, + queryInput: input.queryInput, + answerInput: input.answerInput, + fallbackQuestion: input.fallbackQuestion, + retrievalPlan: input.retrievalPlan, + mode: input.mode, + }), + ), + ), + { concurrency: input.concurrency }, + ), + ) +} + +async function queryRetrievalNamespace(input: { + readonly namespace: string + readonly queryInput: AgenticRetrievalQuery + readonly answerInput: AnswerQuestionInput + readonly fallbackQuestion: string + readonly retrievalPlan: AgenticRetrievalPlan + readonly mode: RetrievalNamespaceQueryMode +}): Promise { + const startedAt = Date.now() + const retrievalQueryParams = buildRetrievalQueryParams({ + input: input.queryInput, + fallbackQuestion: input.fallbackQuestion, + namespace: input.namespace, + useAgentic: input.answerInput.useAgentic ?? true, + sources: input.answerInput.sources, + excludedSourceIds: input.answerInput.excludedSourceIds, + }) + logger.info("chat-agent: searchSources start", { + namespace: input.namespace, + query: retrievalQueryParams.query, + topK: retrievalQueryParams.topK, + useAgentic: retrievalQueryParams.useAgentic, + dataType: retrievalQueryParams.dataType ?? null, + signalPathCount: retrievalQueryParams.signalPaths?.length ?? 0, + filterMode: retrievalQueryParams.filterMode ?? null, + threshold: retrievalQueryParams.threshold ?? null, + targetContent: input.retrievalPlan.targetContent, + purpose: input.retrievalPlan.purpose, + mode: input.mode, + }) + + try { + const response = await input.answerInput.retrieval.query(retrievalQueryParams) + logger.info("chat-agent: searchSources ok", { + namespace: input.namespace, + query: response.query, + durationMs: Date.now() - startedAt, + resultCount: response.results.length, + referencedChunkCount: response.referencedChunks.length, + stopReason: response.stopReason ?? null, + failureReason: response.failureReason ?? null, + targetContent: input.retrievalPlan.targetContent, + mode: input.mode, + }) + logger.info("chat-agent: knowhere query response", { + durationMs: Date.now() - startedAt, + response: formatKnowhereQueryResponseForLog(response), + mode: input.mode, + }) + return { + status: "success", + namespace: input.namespace, + response, + } + } catch (error) { + logger.error("chat-agent: searchSources failed", { + namespace: input.namespace, + query: retrievalQueryParams.query, + durationMs: Date.now() - startedAt, + error: formatUnknownError(error), + targetContent: input.retrievalPlan.targetContent, + mode: input.mode, + rateLimited: isRateLimitError(error), + }) + return { + status: "failure", + namespace: input.namespace, + error, + } + } +} + +function isRateLimitError(error: unknown): boolean { + const statusCode = + getUnknownProperty(error, "statusCode") ?? + getUnknownProperty(error, "status") + if (statusCode === 429 || statusCode === "429") return true + + const body = getUnknownProperty(error, "body") + const bodyError = getUnknownProperty(body, "error") + return [ + getUnknownProperty(error, "name"), + getUnknownProperty(error, "code"), + getUnknownProperty(error, "message"), + getUnknownProperty(body, "message"), + getUnknownProperty(bodyError, "code"), + getUnknownProperty(bodyError, "message"), + ].some( + (value) => + typeof value === "string" && + /\b429\b|rate[\s_-]?limit|too many concurrent|resource_exhausted/i.test( + value, + ), + ) +} + +function getUnknownProperty(value: unknown, key: string): unknown { + if (typeof value !== "object" || value === null) return undefined + try { + return Reflect.get(value, key) + } catch { + return undefined + } +} + function getRetrievalNamespaces(input: AnswerQuestionInput): readonly string[] { const candidates = input.namespaces && input.namespaces.length > 0