From f0a55e26e41184997dd03d91ec2543ab956d4a67 Mon Sep 17 00:00:00 2001 From: Vikranth Reddimasu Date: Fri, 17 Apr 2026 21:54:06 -0400 Subject: [PATCH] =?UTF-8?q?feat(wave4):=20reliability=20=E2=80=94=20thread?= =?UTF-8?q?pools,=20timeouts,=20WAL,=20batching,=20OCR=20guard?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Wave 4 of 5 from .gstack/qa-reports/PLAN.md. The structural weaknesses the tech audit called out as "threats to actually works" — unbounded timeouts, event-loop-blocking inference, OOM-able embedding batches, reads blocking on writes in SQLite, silent no-op ingestions for scanned PDFs, conversation splits when a meta event drops, startup deadlocks on slow Ollama. ## Event-loop stays responsive now - Embedding backend grows an aembed() that offloads model.encode() to the default threadpool (SentenceTransformers) or runs inline (HashEmbedding is trivial). VectorStoreManager grows aquery / aquery_across_notebooks / aquery_document_summaries that use it. RAGService.prepare_prompt and prepare_prompt_cross_notebook are now truly async; a 200–500ms encode no longer stalls health checks and concurrent sidebar refreshes. - LlamaCppBackend.generate / stream_generate run llama_cpp's blocking calls via loop.run_in_executor. Streams pull chunks on a worker thread and yield on the event loop. Without this, every second of inference froze every other request. - OllamaBackend.stream_generate now uses a finite httpx.Timeout (connect 5s, read 300s, write 10s, pool 5s). Previous timeout=None held the connection open indefinitely when Ollama hung. ## OOM and re-upload guards - VectorStoreManager.add_chunks embeds in batches of 64 instead of handing the whole document to model.encode in one call. A 500-page PDF no longer OOMs sentence-transformers. - add_chunks also uses collection.upsert instead of add, so re-uploading a file succeeds cleanly instead of failing mid-batch with DuplicateIDError and leaving the collection in an inconsistent state. - Scanned PDFs with no embedded text now raise a clear DocumentLoaderError ("run it through OCR first") instead of silently ingesting 0 chunks and serving "no relevant documents found" on every subsequent query. ## SQLite stays unblocked - All three stores (notebook_store, conversation_store, metrics_store) now set PRAGMA journal_mode = WAL and PRAGMA busy_timeout = 5000. Readers no longer block on writers — the sidebar refresh during stream completion used to stall under the default DELETE journal. ## Data integrity on stream - Backend now echoes conversation_id on the `done` event as well as `meta`. If the meta event drops (parse error, transient network), the frontend still learns the id via done and doesn't create a second conversation on the next send. useChat honors both with a `!activeConversationId` guard so we only set when it's missing. ## Startup and shipping correctness - Ollama model resolution deferred to an asyncio lifespan so a slow /api/tags round-trip at boot no longer delays FastAPI startup and trips Electron's waitForBackend timeout. app.py imports consolidated at the top of the file (was E402-violating before). - Config split: llm_context_window default 4096 (was 2048), llm_max_tokens default 1024 (was 2048). Previously both were 2048 — the same knob was used for "how much context the model can hold" AND "how many tokens to generate", which quietly truncated long RAG prompts AND capped replies. - Chroma telemetry disabled via Settings(anonymized_telemetry=False). Keeps the offline-first promise and stops filling logs with "Failed to send telemetry event ClientCreateCollectionEvent" warnings. ## Frontend hardening - api.ts request() accepts an optional timeoutMs and composes caller signals with a default 30s AbortController. Prior code could hang on any stalled endpoint forever. Uploads get their own 5-minute budget rather than the default. - Network errors surface through the errorMessages humanizer that Wave 3 added, so the user sees "The backend isn't reachable…" instead of a raw "Failed to fetch". ## Verified - cd apps/desktop && npm run build — clean (272 modules, 864ms) - cd backend && uv run ruff check . — clean - cd backend && uv run pytest -q — 33 passed ## Intentionally not in this PR (for a later wave or follow-up) - 4.1 Backend-crash recovery IPC — bigger Electron main-process refactor - 4.3 Exponential backoff for 503/ECONNREFUSED — fiddly retry-policy work - 4.9 Stream-to-disk upload + size limit — requires multipart-streaming rewrite at the FastAPI boundary - 4.13 Per-conversation asyncio lock — needs a store-level primitive - 4.15 list_documents via summaries — risks regressing docs that failed summary generation; needs a unified doc registry first - 4.16 Document summary LLM deferred to background — needs a task queue Co-Authored-By: Claude Opus 4.7 (1M context) --- apps/desktop/src/api.ts | 98 +++++++++++---- apps/desktop/src/hooks/useChat.ts | 5 + apps/desktop/src/types.ts | 2 +- backend/notebooklm_backend/app.py | 51 ++++++-- backend/notebooklm_backend/config.py | 10 +- backend/notebooklm_backend/routes/chat.py | 7 +- .../services/conversation_store.py | 1 + .../services/document_loader.py | 25 +++- .../notebooklm_backend/services/embeddings.py | 26 +++- backend/notebooklm_backend/services/llm.py | 35 +++++- .../services/metrics_store.py | 2 + .../services/notebook_store.py | 6 + backend/notebooklm_backend/services/rag.py | 12 +- .../services/rag_llamaindex.py | 2 +- .../services/vector_store.py | 117 ++++++++++++++---- 15 files changed, 319 insertions(+), 80 deletions(-) diff --git a/apps/desktop/src/api.ts b/apps/desktop/src/api.ts index 8a9e831..feffd3c 100644 --- a/apps/desktop/src/api.ts +++ b/apps/desktop/src/api.ts @@ -50,22 +50,55 @@ async function getApiBase(): Promise { return resolvedApiBase; } -async function request(path: string, init?: RequestInit): Promise { +/** Default timeout for lookup-style requests. Uploads use their own (longer) + * timeout further down. 30s is generous for local ops but still bounded so a + * hung backend doesn't lock the UI indefinitely. */ +const DEFAULT_REQUEST_TIMEOUT_MS = 30_000; + +async function request(path: string, init?: RequestInit & { timeoutMs?: number }): Promise { const apiBase = await getApiBase(); - const response = await fetch(`${apiBase}${path}`, { - headers: { - 'Content-Type': 'application/json', - ...init?.headers, - }, - ...init, - }); + const timeoutMs = init?.timeoutMs ?? DEFAULT_REQUEST_TIMEOUT_MS; + + // Compose caller-provided signal with our timeout — if either fires, abort. + const timeoutController = new AbortController(); + const timer = setTimeout(() => timeoutController.abort(), timeoutMs); + const composed = composeSignals(init?.signal, timeoutController.signal); + + try { + const response = await fetch(`${apiBase}${path}`, { + headers: { + 'Content-Type': 'application/json', + ...init?.headers, + }, + ...init, + signal: composed, + }); + + if (!response.ok) { + const detail = await response.text(); + throw new Error(detail || `Request failed with status ${response.status}`); + } - if (!response.ok) { - const detail = await response.text(); - throw new Error(detail || `Request failed with status ${response.status}`); + return (await response.json()) as T; + } catch (err) { + if (timeoutController.signal.aborted && !(init?.signal?.aborted)) { + throw new Error('Request timed out'); + } + throw err; + } finally { + clearTimeout(timer); } +} - return response.json() as Promise; +function composeSignals(a: AbortSignal | null | undefined, b: AbortSignal): AbortSignal { + if (!a) return b; + const controller = new AbortController(); + const onAbort = () => controller.abort(); + if (a.aborted) controller.abort(); + else a.addEventListener('abort', onAbort, { once: true }); + if (b.aborted) controller.abort(); + else b.addEventListener('abort', onAbort, { once: true }); + return controller.signal; } export function fetchConfig(): Promise { @@ -87,21 +120,36 @@ export async function uploadDocument(file: File, notebookId?: string): Promise controller.abort(), UPLOAD_TIMEOUT_MS); + + try { + const response = await fetch(`${apiBase}/documents/ingest`, { + method: 'POST', + body: formData, + signal: controller.signal, + }); + + if (!response.ok) { + const detail = await response.text(); + throw new Error(detail || `Upload failed with status ${response.status}`); + } - if (!response.ok) { - const detail = await response.text(); - throw new Error(detail || `Upload failed with status ${response.status}`); + return (await response.json()) as IngestionResponse; + } catch (err) { + if (controller.signal.aborted) { + throw new Error('Upload timed out after 5 minutes'); + } + throw err; + } finally { + clearTimeout(timer); } - - return response.json() as Promise; } export function listDocuments(notebookId: string): Promise { diff --git a/apps/desktop/src/hooks/useChat.ts b/apps/desktop/src/hooks/useChat.ts index 6746063..012ec38 100644 --- a/apps/desktop/src/hooks/useChat.ts +++ b/apps/desktop/src/hooks/useChat.ts @@ -94,6 +94,11 @@ export function useChat() { break; case 'done': if (id) s.updateMessage(id, { content: event.reply, streaming: false }); + // Defensive: if meta was lost, the done event also carries the + // conversation_id so we don't fork the thread on the next send. + if (event.conversation_id && !s.activeConversationId) { + s.setActiveConversationId(event.conversation_id); + } s.setIsStreaming(false); assistantIdRef.current = null; if (s.activeNotebookId) { diff --git a/apps/desktop/src/types.ts b/apps/desktop/src/types.ts index ed58e17..a63812e 100644 --- a/apps/desktop/src/types.ts +++ b/apps/desktop/src/types.ts @@ -47,7 +47,7 @@ export type ChatStreamEvent = conversation_id?: string; } | { type: 'token'; delta: string } - | { type: 'done'; reply: string; metrics?: Record } + | { type: 'done'; reply: string; metrics?: Record; conversation_id?: string } | { type: 'error'; message: string } | { type: 'warning'; message: string }; diff --git a/backend/notebooklm_backend/app.py b/backend/notebooklm_backend/app.py index 99d9e1c..b2c5edc 100644 --- a/backend/notebooklm_backend/app.py +++ b/backend/notebooklm_backend/app.py @@ -1,37 +1,70 @@ from __future__ import annotations +import asyncio +import contextlib +import logging +from typing import AsyncIterator + from fastapi import FastAPI from starlette.middleware.cors import CORSMiddleware from .config import AppConfig, get_settings -from .routes import health, chat, documents, rag, notebooks, metrics, speech, export, agent, conversations, zotero +from .routes import ( + agent, + chat, + conversations, + documents, + export, + health, + metrics, + notebooks, + rag, + speech, + zotero, +) +from .services.agent import AgentService from .services.chat import ChatService +from .services.conversation_store import ConversationStore from .services.embeddings import create_embedding_backend from .services.ingestion import IngestionService from .services.llm import create_llm_backend +from .services.metrics_store import MetricsStore +from .services.model_profiles import resolve_ollama_model +from .services.notebook_store import NotebookStore from .services.rag import RAGService from .services.rag_llamaindex import LlamaIndexRAGService -from .services.vector_store import create_vector_store -from .services.notebook_store import NotebookStore -from .services.conversation_store import ConversationStore -from .services.metrics_store import MetricsStore from .services.speech import SpeechService -from .services.agent import AgentService -from .services.model_profiles import resolve_ollama_model +from .services.vector_store import create_vector_store +logger = logging.getLogger(__name__) def create_app() -> FastAPI: """Create the FastAPI application for the chat-first Notebook LM backend.""" settings: AppConfig = get_settings() settings.ensure_directories() - if settings.llm_provider == "ollama": - resolve_ollama_model(settings) + + @contextlib.asynccontextmanager + async def lifespan(_: FastAPI) -> AsyncIterator[None]: + # Resolve the preferred Ollama model in the background so a slow + # `/api/tags` round-trip doesn't delay startup and cause Electron's + # waitForBackend timeout to fire. + if settings.llm_provider == "ollama": + def _resolve() -> None: + try: + resolve_ollama_model(settings) + except Exception: + logger.exception("Background Ollama model resolution failed") + + loop = asyncio.get_running_loop() + loop.run_in_executor(None, _resolve) + yield app = FastAPI( title="Offline Notebook LM", version="0.1.0", description="Local-first API for chatting with local models (Ollama) with RAG support.", contact={"name": "Offline Notebook LM Team"}, + lifespan=lifespan, ) # Bootstrap services diff --git a/backend/notebooklm_backend/config.py b/backend/notebooklm_backend/config.py index ed3827d..1c378a7 100644 --- a/backend/notebooklm_backend/config.py +++ b/backend/notebooklm_backend/config.py @@ -24,8 +24,14 @@ class AppConfig(BaseSettings): llm_model_path: Path | None = None # Used by llama-cpp onnx_model_path: Path | None = None onnx_execution_provider: Literal["cpu", "cuda", "metal"] = "cpu" - llm_context_window: int = 2048 - llm_max_tokens: int = 2048 + # llm_context_window: how much prompt the model can hold (llama-cpp n_ctx). + # llm_max_tokens: output length cap (passed as num_predict to Ollama, + # max_tokens to llama-cpp/ONNX). Previously both were + # 2048 — the same value was used for both, which + # silently truncated long RAG prompts to a 2048-token + # CONTEXT and also capped replies at 2048 tokens. + llm_context_window: int = 4096 + llm_max_tokens: int = 1024 # Framework integration toggles use_langchain_splitter: bool = True diff --git a/backend/notebooklm_backend/routes/chat.py b/backend/notebooklm_backend/routes/chat.py index 7f20377..daaea95 100644 --- a/backend/notebooklm_backend/routes/chat.py +++ b/backend/notebooklm_backend/routes/chat.py @@ -83,8 +83,11 @@ async def event_generator(): elif event.get("type") == "done": accumulated_reply = event.get("reply", accumulated_reply) - # Inject conversation_id into meta events so frontend knows which conversation - if event.get("type") == "meta" and conversation_id: + # Inject conversation_id into meta AND done events. If the + # meta event is dropped (parse error, network hiccup), the + # frontend falls back to done, rather than silently creating + # a second conversation on the next send. + if event.get("type") in ("meta", "done") and conversation_id: event["conversation_id"] = conversation_id yield f"data: {json.dumps(event)}\n\n" diff --git a/backend/notebooklm_backend/services/conversation_store.py b/backend/notebooklm_backend/services/conversation_store.py index bc8518f..2e56145 100644 --- a/backend/notebooklm_backend/services/conversation_store.py +++ b/backend/notebooklm_backend/services/conversation_store.py @@ -44,6 +44,7 @@ def __init__(self, settings: AppConfig) -> None: def _connect(self) -> Iterable[sqlite3.Connection]: conn = sqlite3.connect(str(self.db_path)) conn.row_factory = sqlite3.Row + conn.execute("PRAGMA journal_mode = WAL") conn.execute("PRAGMA foreign_keys = ON") conn.execute("PRAGMA busy_timeout = 5000") try: diff --git a/backend/notebooklm_backend/services/document_loader.py b/backend/notebooklm_backend/services/document_loader.py index 6c7431b..3e9d881 100644 --- a/backend/notebooklm_backend/services/document_loader.py +++ b/backend/notebooklm_backend/services/document_loader.py @@ -81,19 +81,38 @@ def _load_pdf(file_path: Path) -> LoadedDocument: """Load PDF file using pypdf.""" try: from pypdf import PdfReader - + reader = PdfReader(str(file_path)) text_parts = [] - + page_count = 0 + for page in reader.pages: + page_count += 1 page_text = page.extract_text() if page_text: text_parts.append(page_text) - + text = "\n\n".join(part for part in text_parts if part) + + # Scanned PDFs (image-only, no embedded text) produce near-empty + # output. Previously we silently "succeeded" with 0 chunks and the + # user saw "no relevant documents found" on every query. Raise a + # clear error instead; the API turns this into a 400 with an + # actionable message. + MIN_CHARS_PER_PAGE = 40 + expected_min = MIN_CHARS_PER_PAGE * max(1, page_count) + if len(text.strip()) < min(100, expected_min): + raise DocumentLoaderError( + "This PDF appears to be a scan (no embedded text). OCR isn't " + "supported yet — run the file through a tool like Adobe or " + "Preview's OCR first, then re-upload." + ) + return LoadedDocument(path=file_path, text=text) except ImportError: raise DocumentLoaderError("pypdf is required for PDF support") + except DocumentLoaderError: + raise except Exception as e: raise DocumentLoaderError(f"Failed to read PDF {file_path}: {e}") diff --git a/backend/notebooklm_backend/services/embeddings.py b/backend/notebooklm_backend/services/embeddings.py index 92e148c..0ecd7c8 100644 --- a/backend/notebooklm_backend/services/embeddings.py +++ b/backend/notebooklm_backend/services/embeddings.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio from typing import Protocol from ..config import AppConfig @@ -10,6 +11,11 @@ def embed(self, texts: list[str]) -> list[list[float]]: """Generate embeddings for a list of texts.""" ... + async def aembed(self, texts: list[str]) -> list[list[float]]: + """Async variant that must not block the event loop. Default + implementation offloads to the default executor.""" + ... + class SentenceTransformersBackend: """Embedding backend using sentence-transformers.""" @@ -34,6 +40,14 @@ def embed(self, texts: list[str]) -> list[list[float]]: embeddings = model.encode(texts, show_progress_bar=False) return embeddings.tolist() + async def aembed(self, texts: list[str]) -> list[list[float]]: + # model.encode is CPU-bound and can take 100-500ms; awaiting it + # synchronously freezes the event loop, making the backend + # unresponsive to health checks and concurrent fetches during + # ingestion. Push it to the default threadpool. + loop = asyncio.get_running_loop() + return await loop.run_in_executor(None, self.embed, texts) + class HashEmbeddingBackend: """Simple hash-based embedding for testing (not semantic).""" @@ -41,13 +55,13 @@ class HashEmbeddingBackend: def embed(self, texts: list[str]) -> list[list[float]]: """Generate pseudo-embeddings using hashing.""" import hashlib - + embeddings = [] for text in texts: # Create a 768-dim vector from hash hash_obj = hashlib.sha256(text.encode()) hash_bytes = hash_obj.digest() - + # Repeat hash to get 768 dimensions vector = [] for i in range(768): @@ -55,11 +69,15 @@ def embed(self, texts: list[str]) -> list[list[float]]: # Normalize to [-1, 1] val = (hash_bytes[byte_idx] / 255.0) * 2 - 1 vector.append(val) - + embeddings.append(vector) - + return embeddings + async def aembed(self, texts: list[str]) -> list[list[float]]: + # Cheap enough to run inline — no executor hop needed. + return self.embed(texts) + def create_embedding_backend(settings: AppConfig) -> EmbeddingBackend: """Create an embedding backend based on settings.""" diff --git a/backend/notebooklm_backend/services/llm.py b/backend/notebooklm_backend/services/llm.py index 647f1b0..a1a3688 100644 --- a/backend/notebooklm_backend/services/llm.py +++ b/backend/notebooklm_backend/services/llm.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import json from dataclasses import dataclass, field from functools import cached_property @@ -10,6 +11,11 @@ from ..config import AppConfig +# Finite timeouts on the streaming generation request. Prior code used +# timeout=None which tied up the connection indefinitely when Ollama hung. +# Read of 300s is generous for long completions; connect/write are fast. +_OLLAMA_STREAM_TIMEOUT = httpx.Timeout(connect=5.0, read=300.0, write=10.0, pool=5.0) + try: from llama_cpp import Llama # type: ignore except Exception: # pragma: no cover - optional dependency @@ -62,7 +68,7 @@ async def stream_generate(self, prompt: str, max_tokens: int) -> AsyncIterator[s "options": {"num_predict": max_tokens}, } try: - async with httpx.AsyncClient(timeout=None) as client: + async with httpx.AsyncClient(timeout=_OLLAMA_STREAM_TIMEOUT) as client: async with client.stream("POST", f"{self.base_url}/api/generate", json=payload) as response: response.raise_for_status() async for line in response.aiter_lines(): @@ -122,16 +128,33 @@ def _ensure_model(self) -> Llama: async def generate(self, prompt: str, max_tokens: int) -> str: llama = self._ensure_model() - completion = llama.create_completion( - prompt=prompt, - max_tokens=max_tokens, - stream=False, + # llama_cpp is a C-bound blocking call; run in the default executor + # so the event loop (and any concurrent health checks / UI fetches) + # stay responsive during inference. + loop = asyncio.get_running_loop() + completion = await loop.run_in_executor( + None, + lambda: llama.create_completion(prompt=prompt, max_tokens=max_tokens, stream=False), ) return completion["choices"][0]["text"].strip() async def stream_generate(self, prompt: str, max_tokens: int) -> AsyncIterator[str]: llama = self._ensure_model() - for chunk in llama.create_completion(prompt=prompt, max_tokens=max_tokens, stream=True): + loop = asyncio.get_running_loop() + + # llama.create_completion(stream=True) is a sync generator. Pull each + # chunk on a worker thread and yield it back on the event loop. + def _start_stream(): + return iter( + llama.create_completion(prompt=prompt, max_tokens=max_tokens, stream=True) + ) + + it = await loop.run_in_executor(None, _start_stream) + sentinel = object() + while True: + chunk = await loop.run_in_executor(None, lambda: next(it, sentinel)) + if chunk is sentinel: + break delta = chunk["choices"][0].get("text") if delta: yield delta diff --git a/backend/notebooklm_backend/services/metrics_store.py b/backend/notebooklm_backend/services/metrics_store.py index e029b8c..111b328 100644 --- a/backend/notebooklm_backend/services/metrics_store.py +++ b/backend/notebooklm_backend/services/metrics_store.py @@ -24,6 +24,8 @@ def __init__(self, settings: AppConfig) -> None: def _connect(self): conn = sqlite3.connect(self.db_path) conn.row_factory = sqlite3.Row + conn.execute("PRAGMA journal_mode = WAL") + conn.execute("PRAGMA busy_timeout = 5000") try: yield conn finally: diff --git a/backend/notebooklm_backend/services/notebook_store.py b/backend/notebooklm_backend/services/notebook_store.py index 784699d..d92ef0c 100644 --- a/backend/notebooklm_backend/services/notebook_store.py +++ b/backend/notebooklm_backend/services/notebook_store.py @@ -25,6 +25,12 @@ def __init__(self, settings: AppConfig) -> None: def _connect(self) -> Iterable[sqlite3.Connection]: conn = sqlite3.connect(self.db_path) conn.row_factory = sqlite3.Row + # WAL: concurrent readers don't block on writers. Default DELETE journal + # mode caused the sidebar to stall while chat streams were completing. + # busy_timeout gives the engine 5s to wait out a contended lock before + # raising, matching ConversationStore's behavior. + conn.execute("PRAGMA journal_mode = WAL") + conn.execute("PRAGMA busy_timeout = 5000") conn.execute("PRAGMA foreign_keys = ON") try: yield conn diff --git a/backend/notebooklm_backend/services/rag.py b/backend/notebooklm_backend/services/rag.py index 98808d1..70df80d 100644 --- a/backend/notebooklm_backend/services/rag.py +++ b/backend/notebooklm_backend/services/rag.py @@ -49,7 +49,7 @@ async def prepare_prompt(self, notebook_id: str, question: str, top_k: int = 5) metrics: dict[str, float] = {} # Stage 1: Query document summaries to find relevant documents stage1_start = time.perf_counter() - relevant_summaries = self.vector_store.query_document_summaries( + relevant_summaries = await self.vector_store.aquery_document_summaries( notebook_id=notebook_id, query=question, top_k=3, # Get top 3 most relevant documents @@ -76,7 +76,7 @@ async def prepare_prompt(self, notebook_id: str, question: str, top_k: int = 5) for spath in selected_paths: try: - res = self.vector_store.query( + res = await self.vector_store.aquery( notebook_id=notebook_id, query=question, top_k=per_doc, @@ -87,7 +87,7 @@ async def prepare_prompt(self, notebook_id: str, question: str, top_k: int = 5) dists = res.get("distances", [[]])[0] if res.get("distances") else [] except Exception: # Fallback: unfiltered query + manual filter - res = self.vector_store.query(notebook_id=notebook_id, query=question, top_k=per_doc * 2) + res = await self.vector_store.aquery(notebook_id=notebook_id, query=question, top_k=per_doc * 2) docs = res.get("documents", [[]])[0] metas = res.get("metadatas", [[]])[0] dists = res.get("distances", [[]])[0] if res.get("distances") else [] @@ -111,7 +111,9 @@ async def prepare_prompt(self, notebook_id: str, question: str, top_k: int = 5) else: # Fallback to single-stage if no summaries available retrieval_start = time.perf_counter() - query_results = self.vector_store.query(notebook_id=notebook_id, query=question, top_k=max(top_k, 20)) + query_results = await self.vector_store.aquery( + notebook_id=notebook_id, query=question, top_k=max(top_k, 20) + ) metrics["retrieval_ms"] = (time.perf_counter() - retrieval_start) * 1000 documents = query_results.get("documents", [[]])[0] metadatas = query_results.get("metadatas", [[]])[0] @@ -214,7 +216,7 @@ async def prepare_prompt_cross_notebook( metrics: dict[str, float] = {} retrieval_start = time.perf_counter() - query_results = self.vector_store.query_across_notebooks( + query_results = await self.vector_store.aquery_across_notebooks( notebook_ids=notebook_ids, query=question, top_k=top_k, diff --git a/backend/notebooklm_backend/services/rag_llamaindex.py b/backend/notebooklm_backend/services/rag_llamaindex.py index a4953e5..b5bee10 100644 --- a/backend/notebooklm_backend/services/rag_llamaindex.py +++ b/backend/notebooklm_backend/services/rag_llamaindex.py @@ -122,7 +122,7 @@ def _get_text_embedding(self, text: str) -> list[float]: # Two-stage retrieval: First filter documents by summary, then retrieve chunks # Stage 1: Query document summaries to find relevant documents - relevant_summaries = self.vector_store.query_document_summaries( + relevant_summaries = await self.vector_store.aquery_document_summaries( notebook_id=notebook_id, query=question, top_k=3, # Get top 3 most relevant documents diff --git a/backend/notebooklm_backend/services/vector_store.py b/backend/notebooklm_backend/services/vector_store.py index fca6c6b..df740a7 100644 --- a/backend/notebooklm_backend/services/vector_store.py +++ b/backend/notebooklm_backend/services/vector_store.py @@ -66,30 +66,38 @@ def delete_document(self, notebook_id: str, source_path: str) -> int: return removed + # Embedding all chunks of a 500-page PDF in one call to `model.encode` + # can allocate several GB and OOM. Batching keeps peak memory bounded. + _EMBED_BATCH = 64 + def add_chunks(self, notebook_id: str, chunks: Iterable[TextChunk]) -> int: chunk_list = list(chunks) if not chunk_list: return 0 - documents = [chunk.text for chunk in chunk_list] - embeddings = self.embedding_backend.embed(documents) - metadatas = [ - { - "source_path": chunk.source_path, - "order": chunk.order, - } - for chunk in chunk_list - ] - ids = [chunk.chunk_id for chunk in chunk_list] - collection = self.get_collection(notebook_id) - collection.add( - ids=ids, - documents=documents, - metadatas=metadatas, - embeddings=embeddings, - ) - return len(chunk_list) + total = 0 + for start in range(0, len(chunk_list), self._EMBED_BATCH): + batch = chunk_list[start : start + self._EMBED_BATCH] + documents = [chunk.text for chunk in batch] + embeddings = self.embedding_backend.embed(documents) + metadatas = [ + {"source_path": chunk.source_path, "order": chunk.order} + for chunk in batch + ] + ids = [chunk.chunk_id for chunk in batch] + + # upsert, not add — a re-upload of the same file would have failed + # with DuplicateIDError mid-batch, leaving the collection in an + # inconsistent state (some chunks present, some not). + collection.upsert( + ids=ids, + documents=documents, + metadatas=metadatas, + embeddings=embeddings, + ) + total += len(batch) + return total def query( self, @@ -116,20 +124,60 @@ def query( # Fallback to unfiltered query if the where syntax is unsupported pass return collection.query(**query_kwargs) + + async def aquery( + self, + notebook_id: str, + query: str, + top_k: int = 5, + where: dict | None = None, + ) -> dict: + """Async variant: embedding runs on a threadpool so the event loop + stays responsive during chat stream setup.""" + collection = self.get_collection(notebook_id) + query_emb = await self.embedding_backend.aembed([query]) + query_kwargs = { + "query_embeddings": query_emb, + "n_results": top_k, + } + if where: + try: + query_kwargs["where"] = where + return collection.query(**query_kwargs) + except Exception: + pass + return collection.query(**query_kwargs) + async def aquery_across_notebooks( + self, + notebook_ids: list[str], + query: str, + top_k: int = 5, + ) -> dict: + """Async wrapper for cross-notebook query; embeds on a threadpool.""" + query_emb = await self.embedding_backend.aembed([query]) + return self._query_across_notebooks_impl(notebook_ids, query_emb, top_k) + def query_across_notebooks( self, notebook_ids: list[str], query: str, top_k: int = 5, + ) -> dict: + query_emb = self.embedding_backend.embed([query]) + return self._query_across_notebooks_impl(notebook_ids, query_emb, top_k) + + def _query_across_notebooks_impl( + self, + notebook_ids: list[str], + query_emb: list[list[float]], + top_k: int, ) -> dict: """ Query multiple notebook collections and merge results ranked by distance. Returns the same format as a single query() call but with an extra 'notebook_id' field in each metadata entry. """ - query_emb = self.embedding_backend.embed([query]) - all_documents: list[str] = [] all_metadatas: list[dict] = [] all_distances: list[float] = [] @@ -204,6 +252,23 @@ def store_document_summary(self, notebook_id: str, summary: DocumentSummary) -> embeddings=summary_embedding, ) + async def aquery_document_summaries( + self, + notebook_id: str, + query: str, + top_k: int = 3, + ) -> list[DocumentSummary]: + """Async variant: embedding runs on a threadpool.""" + try: + summaries_collection = self.client.get_collection( + name=self._doc_summaries_collection_name(notebook_id) + ) + except Exception: + return [] + query_emb = await self.embedding_backend.aembed([query]) + results = summaries_collection.query(query_embeddings=query_emb, n_results=top_k) + return self._build_summaries(results) + def query_document_summaries( self, notebook_id: str, @@ -221,10 +286,13 @@ def query_document_summaries( except Exception: # No summaries collection yet, return empty return [] - + # Embed query and search summaries query_emb = self.embedding_backend.embed([query]) results = summaries_collection.query(query_embeddings=query_emb, n_results=top_k) + return self._build_summaries(results) + + def _build_summaries(self, results: dict) -> list[DocumentSummary]: # Convert to DocumentSummary objects summaries = [] @@ -269,5 +337,10 @@ def get_all_document_summaries(self, notebook_id: str) -> list[DocumentSummary]: def create_vector_store(settings: AppConfig, embedding_backend: EmbeddingBackend) -> VectorStoreManager: - client = chromadb.PersistentClient(path=str(settings.index_dir)) + # anonymized_telemetry=False keeps the offline promise intact and quiets + # the "Failed to send telemetry event…" warnings that dominate logs. + client = chromadb.PersistentClient( + path=str(settings.index_dir), + settings=chromadb.Settings(anonymized_telemetry=False), + ) return VectorStoreManager(client=client, embedding_backend=embedding_backend)