From e0aff39a1b5614e5a325028ed9baec4e41e04eed Mon Sep 17 00:00:00 2001 From: Alex Date: Tue, 25 Aug 2026 12:27:09 +0100 Subject: [PATCH] feat: ingestion optimisations --- application/core/settings.py | 3 + application/parser/embedding_pipeline.py | 182 ++++++++-- application/parser/remote/github_loader.py | 338 +++++++++++++++++-- application/worker.py | 102 +++++- tests/parser/file/test_embedding_pipeline.py | 155 ++++++++- tests/parser/remote/test_github_loader.py | 208 +++++++++++- tests/test_worker_utils.py | 76 +++++ tests/worker/test_ingest_checkpoint.py | 87 ++--- 8 files changed, 1037 insertions(+), 114 deletions(-) diff --git a/application/core/settings.py b/application/core/settings.py index c5dc0d2a..ea226b84 100644 --- a/application/core/settings.py +++ b/application/core/settings.py @@ -49,6 +49,9 @@ class Settings(BaseSettings): EMBEDDINGS_BASE_URL: Optional[str] = None # Remote embeddings API URL (OpenAI-compatible) EMBEDDINGS_KEY: Optional[str] = None # api key for embeddings (if using openai, just copy API_KEY) EMBEDDINGS_MAX_INPUT_TOKENS: Optional[int] = None # truncate each remote embed input to N tokens (overflow lost) + EMBEDDINGS_BATCH_SIZE: int = 32 # chunks per embed request during ingest (1 = legacy per-chunk behaviour) + GITHUB_INGEST_MAX_FILE_BYTES: int = 1048576 # skip repo blobs larger than this (0 = no cap) + GITHUB_INGEST_MAX_WORKERS: int = 8 # parallel file fetches per GitHub repo ingest # Optional directory of operator-supplied model YAMLs, loaded after the # built-in catalog under application/core/models/. Later wins on # duplicate model id. See application/core/models/README.md. diff --git a/application/parser/embedding_pipeline.py b/application/parser/embedding_pipeline.py index b6d3a0a6..287cd75d 100755 --- a/application/parser/embedding_pipeline.py +++ b/application/parser/embedding_pipeline.py @@ -39,31 +39,91 @@ def sanitize_content(content: str) -> str: return content.replace('\x00', '') -# Per-chunk inline retry. Aggressive defaults (tries=10, delay=60) blocked +# Fallback when ``settings.EMBEDDINGS_BATCH_SIZE`` is unset or unusable. +DEFAULT_EMBEDDINGS_BATCH_SIZE = 32 + + +def _resolve_batch_size() -> int: + """Read ``EMBEDDINGS_BATCH_SIZE``, falling back to the default. + + Tolerates a missing, ``None`` or non-numeric setting (tests patch + ``settings`` with a ``MagicMock``) so a misconfiguration degrades to + the default batch rather than breaking ingest. + + Returns: + Chunks per embed request, always >= 1. + """ + raw = getattr(settings, "EMBEDDINGS_BATCH_SIZE", None) + # Explicit type check rather than a bare ``int(raw)``: ``int(MagicMock())`` + # succeeds and yields 1, which would silently drop ingest back to the + # per-chunk behaviour this batching replaces. + if isinstance(raw, bool) or not isinstance(raw, (int, str)): + return DEFAULT_EMBEDDINGS_BATCH_SIZE + try: + size = int(raw) + except (TypeError, ValueError): + return DEFAULT_EMBEDDINGS_BATCH_SIZE + return max(1, size) + + +# Per-batch inline retry. Aggressive defaults (tries=10, delay=60) blocked # the loop for up to 9 min per chunk and wedged the heartbeat: lower the # tail so a transient failure fails-fast and the chunk-progress checkpoint # resumes cleanly on next dispatch. @retry(tries=3, delay=5, backoff=2) +def add_texts_to_store_with_retry( + store: Any, docs: List[Any], source_id: str +) -> None: + """Add a batch of documents to the vector store with retry logic. + + One call per batch replaces one call per chunk: the remote embeddings + API accepts a list, and the store writes the whole batch in a single + transaction, so this collapses N HTTP round-trips and N INSERTs into + one of each. Safe to retry — ``add_texts`` embeds before it inserts + and rolls back on failure, so a failed batch writes nothing. + + Args: + store: The vector store object. + docs: The documents to be added. + source_id: Unique identifier for the source. + + Raises: + Exception: If the batch fails after all retry attempts. + """ + if not docs: + return + try: + texts: List[str] = [] + metadatas: List[Any] = [] + for doc in docs: + # Sanitize content to remove NUL characters that cause ingestion failures + doc.page_content = sanitize_content(doc.page_content) + doc.metadata["source_id"] = str(source_id) + texts.append(doc.page_content) + metadatas.append(doc.metadata) + store.add_texts(texts, metadatas=metadatas) + except Exception as e: + logging.error( + f"Failed to add {len(docs)} document(s) with retry: {e}", exc_info=True + ) + raise + + def add_text_to_store_with_retry(store: Any, doc: Any, source_id: str) -> None: - """Add a document's text and metadata to the vector store with retry logic. - + """Add a single document to the vector store with retry logic. + + Thin wrapper over :func:`add_texts_to_store_with_retry`, kept for the + per-chunk fallback in the embed loop and for callers outside it. + Args: store: The vector store object. doc: The document to be added. source_id: Unique identifier for the source. - + Raises: Exception: If document addition fails after all retry attempts. """ - try: - # Sanitize content to remove NUL characters that cause ingestion failures - doc.page_content = sanitize_content(doc.page_content) - - doc.metadata["source_id"] = str(source_id) - store.add_texts([doc.page_content], metadatas=[doc.metadata]) - except Exception as e: - logging.error(f"Failed to add document with retry: {e}", exc_info=True) - raise + add_texts_to_store_with_retry(store, [doc], source_id) def _init_progress_and_resume_index( @@ -115,6 +175,39 @@ def _record_progress(source_id: str, last_index: int, embedded_chunks: int) -> N ) +def _embed_batch_individually( + store: Any, + docs: List[Any], + batch_start: int, + batch_end: int, + source_id: str, +) -> tuple[Exception | None, int | None]: + """Re-run one failed batch a chunk at a time, checkpointing as it goes. + + Called only after a batch raised. Preserves the pre-batching failure + contract: chunks before the offender are embedded and recorded, and the + returned index is the exact chunk that failed rather than the batch head. + + Args: + store: The vector store object. + docs: The full chunk list. + batch_start: Index of the first chunk in the failed batch. + batch_end: Index one past the last chunk in the failed batch. + source_id: Unique identifier for the source. + + Returns: + ``(None, None)`` when every chunk succeeded on its own, else the + exception and the index of the chunk that failed. + """ + for idx in range(batch_start, batch_end): + try: + add_text_to_store_with_retry(store, docs[idx], source_id) + _record_progress(source_id, last_index=idx, embedded_chunks=idx + 1) + except Exception as e: + return e, idx + return None, None + + def assert_index_complete(source_id: str) -> None: """Raise ``EmbeddingPipelineError`` if ``ingest_chunk_progress`` shows a partial embed for ``source_id``. @@ -275,24 +368,31 @@ def embed_and_store_documents( # tripwire still validates ``embedded == total`` afterwards. loop_start = total_docs - # Process and embed documents + # Process and embed documents, one batch per embed request. Progress is + # checkpointed per batch rather than per chunk, so a 3k-chunk ingest + # writes ~90 progress rows instead of 3k — the per-chunk bookkeeping + # (Neon UPDATE + Celery update_state + SSE) dominated wall-clock more + # than the embedding itself. chunk_error: Exception | None = None failed_idx: int | None = None last_published_pct = -1 source_id_str = str(source_id) progress_span = progress_end - progress_start - for idx in tqdm( - range(loop_start, total_docs), + batch_size = _resolve_batch_size() + batch_starts = list(range(loop_start, total_docs, batch_size)) + for batch_start in tqdm( + batch_starts, desc="Embedding 🦖", - unit="docs", - total=total_docs - loop_start, + unit="batch", + total=len(batch_starts), bar_format="{l_bar}{bar}| Time Left: {remaining}", ): - doc = docs[idx] + batch_end = min(batch_start + batch_size, total_docs) + last_idx = batch_end - 1 try: # Map the embed loop into [progress_start, progress_end]. progress = progress_start + int( - ((idx + 1) / total_docs) * progress_span + (batch_end / total_docs) * progress_span ) task_status.update_state(state="PROGRESS", meta={"current": progress}) @@ -307,21 +407,45 @@ def embed_and_store_documents( { "current": progress, "total": total_docs, - "embedded_chunks": idx + 1, + "embedded_chunks": batch_end, "stage": "embedding", }, scope={"kind": "source", "id": source_id_str}, ) last_published_pct = progress - # Add document to vector store - add_text_to_store_with_retry(store, doc, source_id) - _record_progress(source_id, last_index=idx, embedded_chunks=idx + 1) - except Exception as e: - chunk_error = e - failed_idx = idx - logging.error(f"Error embedding document {idx}: {e}", exc_info=True) - logging.info(f"Saving progress at document {idx} out of {total_docs}") + # Add the batch to the vector store + add_texts_to_store_with_retry( + store, docs[batch_start:batch_end], source_id + ) + _record_progress( + source_id, last_index=last_idx, embedded_chunks=batch_end + ) + except Exception as batch_exc: + # The batch failed as a unit and wrote nothing (``add_texts`` + # embeds before it inserts and rolls back). Re-run it one chunk + # at a time so a single poison chunk — an oversized input the + # embeddings server rejects — costs only itself: the chunks + # before it still land and are checkpointed, and ``failed_idx`` + # names the real offender instead of the batch head. + logging.warning( + f"Batch embed failed for chunks {batch_start}-{last_idx} " + f"({batch_exc}); retrying individually to isolate the failure" + ) + chunk_error, failed_idx = _embed_batch_individually( + store, docs, batch_start, batch_end, source_id + ) + if chunk_error is None: + # Every chunk passed on its own — the batch-level failure was + # transient (or a payload-size limit). Progress is recorded. + continue + logging.error( + f"Error embedding document {failed_idx}: {chunk_error}", + exc_info=True, + ) + logging.info( + f"Saving progress at document {failed_idx} out of {total_docs}" + ) try: store.save_local(folder_name) logging.info("Progress saved successfully") diff --git a/application/parser/remote/github_loader.py b/application/parser/remote/github_loader.py index f51ddc18..5c29a1a3 100644 --- a/application/parser/remote/github_loader.py +++ b/application/parser/remote/github_loader.py @@ -1,13 +1,53 @@ import base64 -import requests +import logging +import mimetypes import time -from typing import List, Optional +from concurrent.futures import ThreadPoolExecutor +from typing import Dict, List, Optional, Tuple + +import requests + +from application.core.settings import settings from application.parser.remote.base import BaseRemote from application.parser.schema.base import Document -import mimetypes -from application.core.settings import settings + +logger = logging.getLogger(__name__) + +# Directory names that hold vendored or generated output. Anything under one +# of these is build product, not source: it bloats the index, costs an API +# call each, and answers no question a user would ask of the repo. +SKIP_DIRECTORIES = { + ".git", ".github/workflows/generated", ".idea", ".vscode", + "__pycache__", ".mypy_cache", ".pytest_cache", ".tox", ".venv", "venv", + "node_modules", "bower_components", "vendor", "third_party", "external", + "dist", "build", "out", "target", "bin", "obj", + "site-packages", "coverage", "htmlcov", ".next", ".nuxt", ".svelte-kit", +} + +# Exact filenames that are machine-generated and semantically empty. +SKIP_FILENAMES = { + "package-lock.json", "yarn.lock", "pnpm-lock.yaml", "bun.lockb", + "composer.lock", "gemfile.lock", "cargo.lock", "poetry.lock", + "pipfile.lock", "go.sum", "mix.lock", "podfile.lock", +} + +# Suffixes that mark minified or map output even under a source directory. +SKIP_SUFFIXES = ( + ".min.js", ".min.css", ".map", ".lock", + ".pyc", ".pyo", ".class", ".jar", ".war", ".o", ".so", ".dylib", ".dll", + ".exe", ".wasm", ".bin", ".pdb", +) + class GitHubLoader(BaseRemote): + """Load a GitHub repository's text files as ``Document`` objects. + + Uses the git *tree* API to enumerate the repository in a single request + rather than walking ``/contents/`` directory by directory, filters out + binaries, vendored output and oversized blobs *before* fetching them, and + downloads the survivors in parallel. + """ + def __init__(self): self.access_token = settings.GITHUB_ACCESS_TOKEN self.headers = { @@ -22,15 +62,22 @@ class GitHubLoader(BaseRemote): """Determine if a file is a text file based on extension.""" # Common text file extensions text_extensions = { - '.txt', '.md', '.markdown', '.rst', '.json', '.xml', '.yaml', '.yml', - '.py', '.js', '.ts', '.jsx', '.tsx', '.java', '.c', '.cpp', '.h', '.hpp', - '.cs', '.go', '.rs', '.rb', '.php', '.swift', '.kt', '.scala', - '.html', '.css', '.scss', '.sass', '.less', - '.sh', '.bash', '.zsh', '.fish', - '.sql', '.r', '.m', '.mat', - '.ini', '.cfg', '.conf', '.config', '.env', - '.gitignore', '.dockerignore', '.editorconfig', - '.log', '.csv', '.tsv' + '.txt', '.md', '.markdown', '.rst', '.adoc', '.org', + '.json', '.jsonc', '.xml', '.yaml', '.yml', '.toml', + '.py', '.pyi', '.js', '.mjs', '.cjs', '.ts', '.tsx', '.jsx', + '.vue', '.svelte', '.astro', + '.java', '.c', '.cc', '.cpp', '.h', '.hpp', '.hxx', + '.cs', '.go', '.rs', '.rb', '.php', '.swift', '.kt', '.kts', + '.scala', '.sc', '.clj', '.cljs', '.edn', '.ex', '.exs', '.erl', + '.hs', '.ml', '.mli', '.fs', '.fsx', '.dart', '.lua', '.jl', + '.zig', '.nim', '.v', '.pl', '.pm', '.groovy', '.gradle', + '.html', '.htm', '.css', '.scss', '.sass', '.less', + '.sh', '.bash', '.zsh', '.fish', '.ps1', '.bat', + '.sql', '.graphql', '.gql', '.proto', '.thrift', + '.tf', '.tfvars', '.hcl', '.dockerfile', '.cmake', '.mk', + '.ini', '.cfg', '.conf', '.config', '.properties', + '.gitignore', '.dockerignore', '.editorconfig', '.gitattributes', + '.csv', '.tsv', } # Get file extension @@ -39,6 +86,14 @@ class GitHubLoader(BaseRemote): if file_lower.endswith(ext): return True + # Extension-less files that are conventionally text. + basename = file_lower.rsplit("/", 1)[-1] + if basename in { + "dockerfile", "makefile", "rakefile", "gemfile", "procfile", + "license", "licence", "notice", "authors", "codeowners", "readme", + }: + return True + # Also check MIME type mime_type, _ = mimetypes.guess_type(file_path) if mime_type and (mime_type.startswith("text") or mime_type in ["application/json", "application/xml"]): @@ -46,6 +101,179 @@ class GitHubLoader(BaseRemote): return False + def should_skip_path(self, file_path: str) -> bool: + """Return ``True`` for vendored, generated or binary-by-name paths. + + Applied to the tree listing before any content is fetched, so a + skipped file costs nothing. Complements :meth:`is_text_file`, which + only knows about extensions. + + Args: + file_path: Repo-relative path, e.g. ``"web/dist/app.min.js"``. + + Returns: + ``True`` when the path should not be ingested. + """ + lowered = file_path.lower() + parts = lowered.split("/") + if any(part in SKIP_DIRECTORIES for part in parts[:-1]): + return True + if parts[-1] in SKIP_FILENAMES: + return True + if lowered.endswith(SKIP_SUFFIXES): + return True + # Dot-directories (.git, .cache, ...) never carry documentation. + if any(part.startswith(".") and part not in {".github"} for part in parts[:-1]): + return True + return False + + def _max_file_bytes(self) -> int: + """Resolve the per-blob size cap; ``0`` disables it.""" + raw = getattr(settings, "GITHUB_INGEST_MAX_FILE_BYTES", None) + if isinstance(raw, bool) or not isinstance(raw, (int, str)): + return 1048576 + try: + return max(0, int(raw)) + except (TypeError, ValueError): + return 1048576 + + def _max_workers(self) -> int: + """Resolve the parallel-fetch width, clamped to a sane range.""" + raw = getattr(settings, "GITHUB_INGEST_MAX_WORKERS", None) + if isinstance(raw, bool) or not isinstance(raw, (int, str)): + return 8 + try: + return max(1, min(32, int(raw))) + except (TypeError, ValueError): + return 8 + + @staticmethod + def normalize_repo(repo_url: str) -> str: + """Reduce a user-pasted repo URL to ``owner/name``. + + Strips the scheme/host, a trailing ``.git`` and any trailing slash — + the three shapes that previously 404'd the contents API. + + Args: + repo_url: Anything from ``owner/name`` to + ``https://github.com/owner/name.git/``. + + Returns: + The ``owner/name`` segment. + """ + repo = (repo_url or "").strip() + if not repo: + return "" + + if repo.startswith("git@"): + # git@github.com:owner/name.git + host, _, path = repo[len("git@"):].partition(":") + if host.lower() != "github.com": + return "" + repo = path + elif "://" in repo: + remainder = repo.split("://", 1)[1] + host, _, path = remainder.partition("/") + # Anything not hosted on github.com is not a repo. Previously such + # a URL was pasted straight into the contents path and 404-looped. + if host.lower() not in {"github.com", "www.github.com"}: + return "" + repo = path + elif repo.lower().startswith(("github.com/", "www.github.com/")): + repo = repo.split("/", 1)[1] + + repo = repo.strip("/") + if repo.lower().endswith(".git"): + repo = repo[: -len(".git")] + parts = [p for p in repo.split("/") if p] + if len(parts) < 2: + return "" + owner, name = parts[0], parts[1] + if name.lower().endswith(".git"): + name = name[: -len(".git")] + # A bare host slipping through ("suat-handbook.netlify.app") has no + # owner segment, so require both halves to look like path segments. + if not owner or not name or ":" in owner or ":" in name: + return "" + return f"{owner}/{name}" + + def get_default_branch(self, repo_name: str) -> str: + """Return the repo's default branch, falling back to ``main``. + + The blob URL used for citations was hard-coded to ``main``, which + produced dead source links for every ``master``-default repo. + """ + try: + response = self._make_request(f"https://api.github.com/repos/{repo_name}") + branch = response.json().get("default_branch") + if branch: + return str(branch) + except Exception as e: + logger.warning( + "Could not resolve default branch for %s (%s); assuming 'main'", + repo_name, e, + ) + return "main" + + def fetch_repo_tree( + self, repo_name: str, branch: str + ) -> Tuple[List[Tuple[str, int]], bool]: + """List every blob in the repo with one recursive tree request. + + Replaces the per-directory ``/contents/`` walk, which cost one API + call per directory (720 for a mid-size repo) before a single file + was read. + + Args: + repo_name: ``owner/name``. + branch: Branch or ref to enumerate. + + Returns: + ``(entries, truncated)`` where ``entries`` is a list of + ``(path, size_bytes)`` and ``truncated`` flags a repo too large + for the tree endpoint (caller should fall back to the walk). + """ + url = ( + f"https://api.github.com/repos/{repo_name}/git/trees/" + f"{branch}?recursive=1" + ) + response = self._make_request(url) + payload = response.json() + if isinstance(payload, dict) and "tree" not in payload and "message" in payload: + raise Exception(f"GitHub API error: {payload.get('message')}") + entries = [ + (item.get("path", ""), int(item.get("size") or 0)) + for item in payload.get("tree", []) + if item.get("type") == "blob" + ] + return entries, bool(payload.get("truncated")) + + def select_files(self, entries: List[Tuple[str, int]]) -> List[str]: + """Filter tree entries down to the paths worth fetching. + + Args: + entries: ``(path, size_bytes)`` pairs from :meth:`fetch_repo_tree`. + + Returns: + Repo-relative paths that pass the skip-list, the text-extension + check and the size cap, in tree order. + """ + max_bytes = self._max_file_bytes() + selected: List[str] = [] + skipped_big = 0 + for path, size in entries: + if not path or self.should_skip_path(path) or not self.is_text_file(path): + continue + if max_bytes and size > max_bytes: + skipped_big += 1 + continue + selected.append(path) + if skipped_big: + logger.info( + "Skipped %d file(s) over the %d-byte cap", skipped_big, max_bytes + ) + return selected + def fetch_file_content(self, repo_url: str, file_path: str) -> Optional[str]: """Fetch file content. Returns None if file should be skipped (binary files or empty files).""" url = f"https://api.github.com/repos/{repo_url}/contents/{file_path}" @@ -91,13 +319,18 @@ class GitHubLoader(BaseRemote): remaining = response.headers.get("X-RateLimit-Remaining", "unknown") reset_time = response.headers.get("X-RateLimit-Reset", "unknown") - print(f"GitHub API 403 Error: {error_msg}") - print(f"Rate limit remaining: {remaining}, Reset time: {reset_time}") + logger.warning("GitHub API 403 Error: %s", error_msg) + logger.warning( + "Rate limit remaining: %s, Reset time: %s", remaining, reset_time + ) if "rate limit" in error_msg.lower(): if attempt < max_retries - 1: wait_time = 2 ** attempt # Exponential backoff - print(f"Rate limit hit, waiting {wait_time} seconds before retry...") + logger.warning( + "Rate limit hit, waiting %s seconds before retry...", + wait_time, + ) time.sleep(wait_time) continue @@ -111,12 +344,27 @@ class GitHubLoader(BaseRemote): raise # If we can't parse the response, raise the original error response.raise_for_status() + elif response.status_code == 401 and self.access_token: + # An expired or revoked PAT makes even public repos 401, which + # is strictly worse than not sending one. Retry unauthenticated + # so a stale credential degrades instead of failing the ingest. + logger.warning( + "GitHub rejected the configured token (401); " + "retrying %s unauthenticated", url, + ) + anon = {"Accept": "application/vnd.github.v3+json"} + anon_response = requests.get(url, headers=anon, timeout=100) + if anon_response.status_code == 200: + return anon_response + anon_response.raise_for_status() + return anon_response else: response.raise_for_status() return response def fetch_repo_files(self, repo_url: str, path: str = "") -> List[str]: + """Walk ``/contents/`` recursively (fallback for truncated trees).""" url = f"https://api.github.com/repos/{repo_url}/contents/{path}" response = self._make_request(url) @@ -138,21 +386,69 @@ class GitHubLoader(BaseRemote): files.extend(self.fetch_repo_files(repo_url, item["path"])) return files + def _list_candidate_files(self, repo_name: str, branch: str) -> List[str]: + """Enumerate ingestable paths, preferring the single tree request.""" + try: + entries, truncated = self.fetch_repo_tree(repo_name, branch) + if not truncated: + return self.select_files(entries) + logger.warning( + "Tree for %s is truncated; falling back to the directory walk", + repo_name, + ) + except Exception as e: + logger.warning( + "Tree listing failed for %s (%s); falling back to the " + "directory walk", repo_name, e, + ) + paths = self.fetch_repo_files(repo_name) + return self.select_files([(p, 0) for p in paths]) + def load_data(self, repo_url: str) -> List[Document]: - repo_name = repo_url.split("github.com/")[-1] - files = self.fetch_repo_files(repo_name) + """Load every ingestable text file in ``repo_url`` as a Document.""" + repo_name = self.normalize_repo(repo_url) + if not repo_name or "/" not in repo_name: + raise ValueError( + f"Not a valid GitHub repository: {repo_url!r}. " + "Expected a github.com URL like https://github.com/owner/name." + ) + branch = self.get_default_branch(repo_name) + files = self._list_candidate_files(repo_name, branch) + logger.info( + "Fetching %d file(s) from %s@%s", len(files), repo_name, branch + ) + + # Fetch in parallel: this phase is pure network latency, and it was + # previously one blocking round-trip per file. + contents: Dict[str, Optional[str]] = {} + + def _fetch(file_path: str) -> Tuple[str, Optional[str]]: + try: + return file_path, self.fetch_file_content(repo_name, file_path) + except Exception as e: + # One unreadable file must not sink the whole repo ingest. + logger.warning("Skipping %s: %s", file_path, e) + return file_path, None + + max_workers = min(self._max_workers(), len(files)) or 1 + with ThreadPoolExecutor(max_workers=max_workers) as pool: + for file_path, content in pool.map(_fetch, files): + contents[file_path] = content + documents = [] for file_path in files: - content = self.fetch_file_content(repo_name, file_path) + content = contents.get(file_path) # Skip binary files (content is None) - if content is None: + if not content: continue documents.append(Document( text=content, doc_id=file_path, extra_info={ "title": file_path, - "source": f"https://github.com/{repo_name}/blob/main/{file_path}" + "source": ( + f"https://github.com/{repo_name}/blob/{branch}/{file_path}" + ), } )) return documents diff --git a/application/worker.py b/application/worker.py index f4149c99..1268ca39 100755 --- a/application/worker.py +++ b/application/worker.py @@ -75,6 +75,82 @@ RECURSION_DEPTH = 2 INGEST_HEARTBEAT_INTERVAL_SECONDS = 30 +def count_structure_files(node: dict) -> int: + """Count leaf files in a nested ``directory_structure`` mapping. + + Directories are plain dicts of children; files are dicts carrying a + ``token_count`` key. ``len()`` on the root only sees top-level entries, + which undercounts any repo with subdirectories. + + Args: + node: A ``directory_structure`` mapping (or any subtree of one). + + Returns: + Number of file leaves beneath ``node``. + """ + if not isinstance(node, dict): + return 0 + total = 0 + for value in node.values(): + if isinstance(value, dict): + if "token_count" in value and "size_bytes" in value: + total += 1 + else: + total += count_structure_files(value) + return total + + +def add_file_to_structure( + directory_structure: dict, + file_path: str, + file_type: str, + *, + size_bytes: int, + token_count: int, +) -> None: + """Insert one chunk's stats into a nested ``directory_structure``. + + Callers feed this *chunks*, so a file larger than one chunk arrives + several times. Stats are accumulated rather than overwritten — the + previous assignment kept only the final fragment, which made the + per-file sizes shown in the UI wrong for every multi-chunk file + (a 2.56M-token repo reported 733k). + + Args: + directory_structure: Mapping mutated in place. + file_path: Repo-relative path, e.g. ``"guides/setup.md"``. + file_type: MIME type recorded on first insert. + size_bytes: Byte length of this chunk. + token_count: Token count of this chunk. + + Returns: + None + """ + path_parts = [p for p in file_path.split("/") if p] + if not path_parts: + return + current_level = directory_structure + for part in path_parts[:-1]: + # Intermediate parts are directories + child = current_level.get(part) + if not isinstance(child, dict) or "token_count" in child: + child = {} + current_level[part] = child + current_level = child + + leaf = path_parts[-1] + existing = current_level.get(leaf) + if isinstance(existing, dict) and "token_count" in existing: + existing["size_bytes"] += size_bytes + existing["token_count"] += token_count + else: + current_level[leaf] = { + "type": file_type, + "size_bytes": size_bytes, + "token_count": token_count, + } + + def graph_extraction_key(source_id, updated_at) -> str: """Build the extract_graph idempotency key for a source's current state. @@ -1264,24 +1340,18 @@ def remote_worker( # Build nested directory structure from path # e.g., "guides/setup.md" -> {"guides": {"setup.md": {...}}} - path_parts = file_path.split("/") - current_level = directory_structure - for i, part in enumerate(path_parts): - if i == len(path_parts) - 1: - # Last part is the file - current_level[part] = { - "type": file_type, - "size_bytes": size_bytes, - "token_count": token_count, - } - else: - # Intermediate parts are directories - if part not in current_level: - current_level[part] = {} - current_level = current_level[part] + add_file_to_structure( + directory_structure, file_path, file_type, + size_bytes=size_bytes, token_count=token_count, + ) + # ``len(directory_structure)`` counts only top-level entries, so a + # 1,474-file repo logged "44 files". Count the leaves instead — this + # line is the operational signal that a remote ingest succeeded. logging.info( - f"Built directory structure with {len(directory_structure)} files: " + f"Built directory structure with " + f"{count_structure_files(directory_structure)} files across " + f"{len(directory_structure)} top-level entries: " f"{list(directory_structure.keys())}" ) diff --git a/tests/parser/file/test_embedding_pipeline.py b/tests/parser/file/test_embedding_pipeline.py index 589124e8..19c2d59e 100644 --- a/tests/parser/file/test_embedding_pipeline.py +++ b/tests/parser/file/test_embedding_pipeline.py @@ -3,8 +3,11 @@ import logging from unittest.mock import patch, MagicMock from application.parser.embedding_pipeline import ( + DEFAULT_EMBEDDINGS_BATCH_SIZE, EmbeddingPipelineError, + _resolve_batch_size, add_text_to_store_with_retry, + add_texts_to_store_with_retry, assert_index_complete, embed_and_store_documents, sanitize_content, @@ -123,7 +126,7 @@ def test_embed_and_store_documents_progress_band( assert currents == sorted(currents) -@patch("application.parser.embedding_pipeline.add_text_to_store_with_retry") +@patch("application.parser.embedding_pipeline.add_texts_to_store_with_retry") def test_embed_and_store_documents_partial_failure_raises( mock_add_retry, tmp_path, mock_settings, mock_vector_creator, caplog ): @@ -146,9 +149,11 @@ def test_embed_and_store_documents_partial_failure_raises( mock_vector_creator.create_vectorstore.return_value = mock_store # First document succeeds (FAISS init seeds with docs[0]; the loop - # picks up at idx=1 and raises on the bad chunk). - def side_effect(*args, **kwargs): - if "bad" in args[1].page_content: + # picks up at idx=1 and raises on the bad chunk). The batch entry point + # receives a list, and the per-chunk fallback re-runs it one at a time — + # both go through this mock, so "bad" raises either way. + def side_effect(store_arg, docs_arg, source_arg): + if any("bad" in d.page_content for d in docs_arg): raise RuntimeError("Embedding failed") mock_add_retry.side_effect = side_effect @@ -165,7 +170,7 @@ def test_embed_and_store_documents_partial_failure_raises( mock_store.save_local.assert_called() -@patch("application.parser.embedding_pipeline.add_text_to_store_with_retry") +@patch("application.parser.embedding_pipeline.add_texts_to_store_with_retry") def test_embed_and_store_documents_all_chunks_succeed_no_raise( mock_add_retry, tmp_path, mock_settings, mock_vector_creator, ): @@ -291,3 +296,143 @@ def test_embed_and_store_documents_save_fails_raises_oserror( with pytest.raises(OSError, match="Unable to save vector store"): embed_and_store_documents(docs, str(folder_name), source_id, task_status) + + +# ── batched embed loop ───────────────────────────────────────────────────── + + +def test_add_texts_to_store_with_retry_sends_one_call_per_batch(): + """The batch entry point collapses N chunks into a single add_texts.""" + store = MagicMock() + docs = [MagicMock(page_content=f"c{i}", metadata={}) for i in range(3)] + + add_texts_to_store_with_retry(store, docs, "sid") + + store.add_texts.assert_called_once_with( + ["c0", "c1", "c2"], + metadatas=[{"source_id": "sid"}] * 3, + ) + + +def test_add_texts_to_store_with_retry_sanitizes_and_skips_empty(): + store = MagicMock() + docs = [MagicMock(page_content="a\x00b", metadata={})] + + add_texts_to_store_with_retry(store, docs, "sid") + assert store.add_texts.call_args.args[0] == ["ab"] + + store.reset_mock() + add_texts_to_store_with_retry(store, [], "sid") + store.add_texts.assert_not_called() + + +def test_resolve_batch_size_falls_back_on_bad_setting(monkeypatch): + fake = MagicMock() # attribute access yields a MagicMock, not an int + monkeypatch.setattr("application.parser.embedding_pipeline.settings", fake) + assert _resolve_batch_size() == DEFAULT_EMBEDDINGS_BATCH_SIZE + + fake.EMBEDDINGS_BATCH_SIZE = 0 + assert _resolve_batch_size() == 1 # never below 1 + fake.EMBEDDINGS_BATCH_SIZE = 32 + assert _resolve_batch_size() == 32 + + +def test_embed_loop_batches_chunks(tmp_path, mock_settings, mock_vector_creator): + """70 chunks at batch size 32 => 3 add_texts calls, not 70.""" + mock_settings.VECTOR_STORE = "chromadb" + mock_settings.EMBEDDINGS_BATCH_SIZE = 32 + + docs = [MagicMock(page_content=f"d{i}", metadata={}) for i in range(70)] + store = MagicMock() + mock_vector_creator.create_vectorstore.return_value = store + + with patch("application.parser.embedding_pipeline._record_progress") as rec: + embed_and_store_documents( + docs, str(tmp_path / "s"), "sid", MagicMock(), + ) + + assert store.add_texts.call_count == 3 + assert [len(c.args[0]) for c in store.add_texts.call_args_list] == [32, 32, 6] + # One checkpoint per batch, and the final one accounts for every chunk. + assert rec.call_count == 3 + assert rec.call_args.kwargs == {"last_index": 69, "embedded_chunks": 70} + + +def test_embed_loop_batch_size_one_matches_legacy( + tmp_path, mock_settings, mock_vector_creator +): + """batch_size=1 restores the pre-batching one-call-per-chunk behaviour.""" + mock_settings.VECTOR_STORE = "chromadb" + mock_settings.EMBEDDINGS_BATCH_SIZE = 1 + + docs = [MagicMock(page_content=f"d{i}", metadata={}) for i in range(5)] + store = MagicMock() + mock_vector_creator.create_vectorstore.return_value = store + + embed_and_store_documents(docs, str(tmp_path / "s"), "sid", MagicMock()) + + assert store.add_texts.call_count == 5 + + +def test_poison_chunk_isolated_by_per_chunk_fallback( + tmp_path, mock_settings, mock_vector_creator +): + """A batch failure re-runs individually: good chunks land, and the + reported failure index is the real offender, not the batch head. + + Patches the batch entry point so the ``@retry`` sleeps don't run — the + fallback, not the retry, is what's under test here. + """ + mock_settings.VECTOR_STORE = "chromadb" + mock_settings.EMBEDDINGS_BATCH_SIZE = 32 + + docs = [MagicMock(page_content=f"d{i}", metadata={}) for i in range(10)] + docs[6].page_content = "poison" + mock_vector_creator.create_vectorstore.return_value = MagicMock() + + def fake_add(store, batch, source_id): + if any(d.page_content == "poison" for d in batch): + raise RuntimeError("input too large") + + with patch( + "application.parser.embedding_pipeline.add_texts_to_store_with_retry", + side_effect=fake_add, + ): + with patch("application.parser.embedding_pipeline._record_progress") as rec: + with pytest.raises(EmbeddingPipelineError) as exc: + embed_and_store_documents( + docs, str(tmp_path / "s"), "sid", MagicMock(), + ) + + assert "chunk 6/10" in str(exc.value) + assert isinstance(exc.value.__cause__, RuntimeError) + # Chunks 0-5 were salvaged one at a time and checkpointed. + assert rec.call_args_list[-1].kwargs == {"last_index": 5, "embedded_chunks": 6} + + +def test_batch_only_failure_recovers_via_fallback( + tmp_path, mock_settings, mock_vector_creator +): + """When the *batch* is rejected but each chunk is fine on its own (e.g. a + request-size limit), the fallback completes the ingest without raising.""" + mock_settings.VECTOR_STORE = "chromadb" + mock_settings.EMBEDDINGS_BATCH_SIZE = 32 + + docs = [MagicMock(page_content=f"d{i}", metadata={}) for i in range(4)] + mock_vector_creator.create_vectorstore.return_value = MagicMock() + + seen = [] + + def fake_add(store, batch, source_id): + seen.append(len(batch)) + if len(batch) > 1: + raise RuntimeError("payload too large") + + with patch( + "application.parser.embedding_pipeline.add_texts_to_store_with_retry", + side_effect=fake_add, + ): + embed_and_store_documents(docs, str(tmp_path / "s"), "sid", MagicMock()) + + # One rejected batch of 4, then four successful singles. + assert seen == [4, 1, 1, 1, 1] diff --git a/tests/parser/remote/test_github_loader.py b/tests/parser/remote/test_github_loader.py index 32662279..d6932d20 100644 --- a/tests/parser/remote/test_github_loader.py +++ b/tests/parser/remote/test_github_loader.py @@ -91,9 +91,11 @@ class TestGitHubLoaderLoadData: loader = GitHubLoader() # Stub out network-dependent methods - monkeypatch.setattr(loader, "fetch_repo_files", lambda repo, path="": [ - "README.md", "src/main.py" - ]) + monkeypatch.setattr(loader, "get_default_branch", lambda repo: "main") + monkeypatch.setattr( + loader, "fetch_repo_tree", + lambda repo, branch: ([("README.md", 10), ("src/main.py", 10)], False), + ) def fake_fetch_content(repo, file_path): return f"content for {file_path}" @@ -257,8 +259,10 @@ class TestGitHubLoaderFetchFileContentEdgeCases: class TestGitHubLoaderLoadDataSkipsNone: def test_skips_binary_files(self, monkeypatch): loader = GitHubLoader() + monkeypatch.setattr(loader, "get_default_branch", lambda repo: "main") monkeypatch.setattr( - loader, "fetch_repo_files", lambda repo, path="": ["a.py", "b.png"] + loader, "fetch_repo_tree", + lambda repo, branch: ([("a.py", 10), ("b.png", 10)], False), ) def fake_content(repo, fp): @@ -312,3 +316,199 @@ class TestGitHubLoaderRobustness: mock_get.return_value = make_response({"encoding": "base64", "content": "AAA"}) result = GitHubLoader().fetch_file_content("owner/repo", "bigfile.bin") assert result is None + + +class TestGitHubLoaderNormalizeRepo: + @pytest.mark.parametrize("raw,expected", [ + ("https://github.com/owner/repo", "owner/repo"), + ("https://github.com/owner/repo.git", "owner/repo"), + ("https://github.com/owner/repo/", "owner/repo"), + ("http://github.com/owner/repo.git/", "owner/repo"), + ("owner/repo", "owner/repo"), + ("https://github.com/owner/repo/tree/main/sub", "owner/repo"), + ]) + def test_normalizes(self, raw, expected): + assert GitHubLoader.normalize_repo(raw) == expected + + def test_rejects_non_repo_url(self): + """Regression: a pasted website URL used to be concatenated straight + into the contents path and 404-loop.""" + loader = GitHubLoader() + with pytest.raises(ValueError, match="Not a valid GitHub repository"): + loader.load_data("https://suat-handbook.netlify.app/") + + +class TestGitHubLoaderSkipPaths: + @pytest.mark.parametrize("path", [ + "node_modules/left-pad/index.js", + "web/dist/app.js", + "target/scala-2.13/Foo.class", + "vendor/github.com/pkg/errors/errors.go", + "package-lock.json", + "assets/app.min.js", + "build/output.map", + ".venv/lib/thing.py", + "src/__pycache__/mod.pyc", + ]) + def test_skipped(self, path): + assert GitHubLoader().should_skip_path(path) is True + + @pytest.mark.parametrize("path", [ + "README.md", + "src/main/scala/zio/ZIO.scala", + ".github/workflows/ci.yml", + "docs/guide/setup.md", + "Cargo.toml", + ]) + def test_kept(self, path): + assert GitHubLoader().should_skip_path(path) is False + + +class TestGitHubLoaderIsTextFileAdditions: + @pytest.mark.parametrize("path", [ + "Cargo.toml", "main.go", "app.vue", "schema.proto", + "infra.tf", "Dockerfile", "Makefile", "build.gradle", + ]) + def test_newly_recognised_text(self, path): + assert GitHubLoader().is_text_file(path) is True + + def test_env_files_are_not_ingested(self): + """.env was in the allowlist, so a committed secrets file would be + embedded verbatim into the vector index.""" + loader = GitHubLoader() + assert loader.is_text_file("config/.env") is False + assert loader.is_text_file(".env") is False + + +class TestGitHubLoaderSelectFiles: + def test_applies_size_cap(self, monkeypatch): + loader = GitHubLoader() + monkeypatch.setattr( + "application.parser.remote.github_loader.settings.GITHUB_INGEST_MAX_FILE_BYTES", + 100, raising=False, + ) + entries = [("small.py", 50), ("huge.py", 5000), ("ok.md", 99)] + assert loader.select_files(entries) == ["small.py", "ok.md"] + + def test_zero_cap_disables_limit(self, monkeypatch): + loader = GitHubLoader() + monkeypatch.setattr( + "application.parser.remote.github_loader.settings.GITHUB_INGEST_MAX_FILE_BYTES", + 0, raising=False, + ) + assert loader.select_files([("huge.py", 10**9)]) == ["huge.py"] + + def test_filters_binaries_and_vendored(self): + entries = [ + ("README.md", 10), ("logo.png", 10), + ("node_modules/x/i.js", 10), ("src/a.py", 10), + ] + assert GitHubLoader().select_files(entries) == ["README.md", "src/a.py"] + + +class TestGitHubLoaderTree: + @patch("application.parser.remote.github_loader.requests.get") + def test_single_request_lists_all_blobs(self, mock_get): + mock_get.return_value = make_response({ + "tree": [ + {"path": "README.md", "type": "blob", "size": 12}, + {"path": "src", "type": "tree"}, + {"path": "src/a.py", "type": "blob", "size": 34}, + ], + "truncated": False, + }) + entries, truncated = GitHubLoader().fetch_repo_tree("owner/repo", "main") + + assert entries == [("README.md", 12), ("src/a.py", 34)] + assert truncated is False + # One call for the whole repo, versus one per directory before. + assert mock_get.call_count == 1 + + @patch("application.parser.remote.github_loader.requests.get") + def test_truncated_tree_falls_back_to_walk(self, mock_get, monkeypatch): + loader = GitHubLoader() + mock_get.return_value = make_response({"tree": [], "truncated": True}) + monkeypatch.setattr( + loader, "fetch_repo_files", lambda repo, path="": ["README.md"] + ) + assert loader._list_candidate_files("owner/repo", "main") == ["README.md"] + + +class TestGitHubLoaderDefaultBranch: + @patch("application.parser.remote.github_loader.requests.get") + def test_uses_repo_default_branch(self, mock_get): + mock_get.return_value = make_response({"default_branch": "master"}) + assert GitHubLoader().get_default_branch("owner/repo") == "master" + + @patch("application.parser.remote.github_loader.requests.get") + def test_falls_back_to_main(self, mock_get): + mock_get.side_effect = requests.ConnectionError("boom") + assert GitHubLoader().get_default_branch("owner/repo") == "main" + + def test_citation_url_uses_real_branch(self, monkeypatch): + """Regression: blob/main was hard-coded, so every master-default + repo produced dead source links.""" + loader = GitHubLoader() + monkeypatch.setattr(loader, "get_default_branch", lambda repo: "master") + monkeypatch.setattr( + loader, "fetch_repo_tree", lambda repo, branch: ([("a.py", 5)], False) + ) + monkeypatch.setattr(loader, "fetch_file_content", lambda r, p: "code") + + docs = loader.load_data("https://github.com/owner/repo") + assert docs[0].extra_info["source"] == ( + "https://github.com/owner/repo/blob/master/a.py" + ) + + +class TestGitHubLoaderParallelFetch: + def test_fetches_in_parallel_and_preserves_order(self, monkeypatch): + loader = GitHubLoader() + monkeypatch.setattr(loader, "get_default_branch", lambda repo: "main") + paths = [f"f{i}.py" for i in range(10)] + monkeypatch.setattr( + loader, "fetch_repo_tree", + lambda repo, branch: ([(p, 5) for p in paths], False), + ) + monkeypatch.setattr( + loader, "fetch_file_content", lambda r, p: f"body {p}" + ) + + docs = loader.load_data("https://github.com/owner/repo") + assert [d.doc_id for d in docs] == paths + + def test_one_bad_file_does_not_sink_the_ingest(self, monkeypatch): + loader = GitHubLoader() + monkeypatch.setattr(loader, "get_default_branch", lambda repo: "main") + monkeypatch.setattr( + loader, "fetch_repo_tree", + lambda repo, branch: ([("good.py", 5), ("bad.py", 5)], False), + ) + + def flaky(repo, path): + if path == "bad.py": + raise requests.HTTPError("500") + return "code" + monkeypatch.setattr(loader, "fetch_file_content", flaky) + + docs = loader.load_data("https://github.com/owner/repo") + assert [d.doc_id for d in docs] == ["good.py"] + + +class TestGitHubLoaderStaleTokenFallback: + @patch("application.parser.remote.github_loader.requests.get") + def test_401_retries_unauthenticated(self, mock_get): + """An expired PAT 401s even public repos; fall back to anonymous + rather than failing the ingest outright.""" + loader = GitHubLoader() + loader.access_token = "stale-token" + unauthorized = MagicMock(status_code=401) + ok = make_response({"ok": True}, 200) + mock_get.side_effect = [unauthorized, ok] + + resp = loader._make_request("https://api.github.com/repos/o/r") + + assert resp.status_code == 200 + assert mock_get.call_count == 2 + # Second attempt carried no Authorization header. + assert "Authorization" not in mock_get.call_args.kwargs["headers"] diff --git a/tests/test_worker_utils.py b/tests/test_worker_utils.py index bd6e4021..33ea90cf 100644 --- a/tests/test_worker_utils.py +++ b/tests/test_worker_utils.py @@ -311,3 +311,79 @@ class TestUploadIndex: mock_post.assert_called_once() files = mock_post.call_args.kwargs["files"] assert "file_faiss" in files and "file_pkl" in files + + +# ── directory_structure accounting ───────────────────────────────────────── + + +class TestCountStructureFiles: + """Regression: the remote-ingest completion log used + ``len(directory_structure)``, which counts top-level entries only — a + 1,474-file repo logged "44 files".""" + + def test_counts_nested_leaves_not_top_level_keys(self): + from application.worker import count_structure_files + + structure = { + "README.md": {"type": "text/markdown", "size_bytes": 10, "token_count": 3}, + "src": { + "main": { + "a.py": {"type": "text/x-python", "size_bytes": 5, "token_count": 2}, + "b.py": {"type": "text/x-python", "size_bytes": 5, "token_count": 2}, + }, + "c.py": {"type": "text/x-python", "size_bytes": 5, "token_count": 2}, + }, + } + assert len(structure) == 2 # the old, wrong number + assert count_structure_files(structure) == 4 + + def test_empty_and_non_dict_are_zero(self): + from application.worker import count_structure_files + + assert count_structure_files({}) == 0 + assert count_structure_files(None) == 0 + assert count_structure_files({"empty_dir": {}}) == 0 + + +class TestDirectoryStructureAccumulatesChunks: + """Regression: ``directory_structure`` was built from *chunks* but + assigned per path, so a multi-chunk file kept only its last fragment's + size and token count.""" + + def test_multi_chunk_file_sums_its_chunks(self): + from application.worker import add_file_to_structure, count_structure_files + + structure: dict = {} + for tokens, size in [(900, 3600), (900, 3600), (120, 480)]: + add_file_to_structure( + structure, "docs/guide.md", "text/markdown", + size_bytes=size, token_count=tokens, + ) + + leaf = structure["docs"]["guide.md"] + assert leaf["token_count"] == 1920 # not 120, the last chunk + assert leaf["size_bytes"] == 7680 # not 480 + assert leaf["type"] == "text/markdown" + assert count_structure_files(structure) == 1 + + def test_distinct_files_stay_separate(self): + from application.worker import add_file_to_structure, count_structure_files + + structure: dict = {} + for path in ["a.md", "src/b.py", "src/deep/c.py"]: + add_file_to_structure( + structure, path, "text/plain", size_bytes=10, token_count=4, + ) + + assert count_structure_files(structure) == 3 + assert structure["src"]["deep"]["c.py"]["token_count"] == 4 + assert structure["a.md"]["token_count"] == 4 + + def test_empty_path_is_ignored(self): + from application.worker import add_file_to_structure + + structure: dict = {} + add_file_to_structure( + structure, "", "text/plain", size_bytes=1, token_count=1, + ) + assert structure == {} diff --git a/tests/worker/test_ingest_checkpoint.py b/tests/worker/test_ingest_checkpoint.py index 7313433d..347898da 100644 --- a/tests/worker/test_ingest_checkpoint.py +++ b/tests/worker/test_ingest_checkpoint.py @@ -89,21 +89,22 @@ class TestEmbedCheckpoint: docs = _make_docs(5) captured_chunks: list[str] = [] - def _fake_add(store, doc, sid): - captured_chunks.append(doc.page_content) + def _fake_add(store, batch, sid): + captured_chunks.extend(d.page_content for d in batch) # The retry decorator wraps the real fn. Patch the module-level name - # so our loop calls the spy directly. + # so our loop calls the spy directly. The loop embeds a batch at a + # time, so the spy receives a list and flattens it back out. import application.parser.embedding_pipeline as ep_mod - original = ep_mod.add_text_to_store_with_retry - ep_mod.add_text_to_store_with_retry = _fake_add + original = ep_mod.add_texts_to_store_with_retry + ep_mod.add_texts_to_store_with_retry = _fake_add try: ep_mod.embed_and_store_documents( docs, str(tmp_path), source_id, MagicMock() ) finally: - ep_mod.add_text_to_store_with_retry = original + ep_mod.add_texts_to_store_with_retry = original # Faiss seeds the store with docs[0]; the loop picks up at idx=1. assert captured_chunks == ["chunk-1", "chunk-2", "chunk-3", "chunk-4"] @@ -136,17 +137,17 @@ class TestEmbedCheckpoint: captured_chunks: list[str] = [] - def _fake_add(store, doc, sid): - captured_chunks.append(doc.page_content) + def _fake_add(store, batch, sid): + captured_chunks.extend(d.page_content for d in batch) - original = ep_mod.add_text_to_store_with_retry - ep_mod.add_text_to_store_with_retry = _fake_add + original = ep_mod.add_texts_to_store_with_retry + ep_mod.add_texts_to_store_with_retry = _fake_add try: ep_mod.embed_and_store_documents( docs, str(tmp_path), source_id, MagicMock() ) finally: - ep_mod.add_text_to_store_with_retry = original + ep_mod.add_texts_to_store_with_retry = original # On resume the FAISS store is loaded from storage (no docs_init); # the loop iterates the un-popped docs list starting at resume_index. @@ -187,14 +188,14 @@ class TestEmbedCheckpoint: docs = _make_docs(6) _seed_progress_row(pg_conn, source_id, total=6, last_index=2) - original = ep_mod.add_text_to_store_with_retry - ep_mod.add_text_to_store_with_retry = lambda store, doc, sid: None + original = ep_mod.add_texts_to_store_with_retry + ep_mod.add_texts_to_store_with_retry = lambda store, batch, sid: None try: ep_mod.embed_and_store_documents( docs, str(tmp_path), source_id, MagicMock() ) finally: - ep_mod.add_text_to_store_with_retry = original + ep_mod.add_texts_to_store_with_retry = original assert len(captured_kwargs) == 1 # No ``docs_init`` on resume — this is what triggers FaissStore to @@ -219,14 +220,14 @@ class TestEmbedCheckpoint: docs = _make_docs(4) _seed_progress_row(pg_conn, source_id, total=4, last_index=1) - original = ep_mod.add_text_to_store_with_retry - ep_mod.add_text_to_store_with_retry = lambda store, doc, sid: None + original = ep_mod.add_texts_to_store_with_retry + ep_mod.add_texts_to_store_with_retry = lambda store, batch, sid: None try: ep_mod.embed_and_store_documents( docs, str(tmp_path), source_id, MagicMock() ) finally: - ep_mod.add_text_to_store_with_retry = original + ep_mod.add_texts_to_store_with_retry = original fake_store.delete_index.assert_not_called() @@ -245,9 +246,11 @@ class TestEmbedCheckpoint: ) captured: list[str] = [] - original = ep_mod.add_text_to_store_with_retry - ep_mod.add_text_to_store_with_retry = ( - lambda store, doc, sid: captured.append(doc.page_content) + original = ep_mod.add_texts_to_store_with_retry + ep_mod.add_texts_to_store_with_retry = ( + lambda store, batch, sid: captured.extend( + d.page_content for d in batch + ) ) try: ep_mod.embed_and_store_documents( @@ -255,7 +258,7 @@ class TestEmbedCheckpoint: attempt_id="att-A", ) finally: - ep_mod.add_text_to_store_with_retry = original + ep_mod.add_texts_to_store_with_retry = original # Same attempt → resume past the last persisted index. assert captured == ["chunk-3", "chunk-4", "chunk-5"] @@ -278,9 +281,11 @@ class TestEmbedCheckpoint: docs = _make_docs(5) captured: list[str] = [] - original = ep_mod.add_text_to_store_with_retry - ep_mod.add_text_to_store_with_retry = ( - lambda store, doc, sid: captured.append(doc.page_content) + original = ep_mod.add_texts_to_store_with_retry + ep_mod.add_texts_to_store_with_retry = ( + lambda store, batch, sid: captured.extend( + d.page_content for d in batch + ) ) try: ep_mod.embed_and_store_documents( @@ -288,7 +293,7 @@ class TestEmbedCheckpoint: attempt_id="att-new", ) finally: - ep_mod.add_text_to_store_with_retry = original + ep_mod.add_texts_to_store_with_retry = original # Fresh attempt → reset to chunk 0; FAISS branch seeds with # docs[0] and the loop picks up at idx=1. @@ -327,9 +332,11 @@ class TestEmbedCheckpoint: docs = _make_docs(5) captured: list[str] = [] - original = ep_mod.add_text_to_store_with_retry - ep_mod.add_text_to_store_with_retry = ( - lambda store, doc, sid: captured.append(doc.page_content) + original = ep_mod.add_texts_to_store_with_retry + ep_mod.add_texts_to_store_with_retry = ( + lambda store, batch, sid: captured.extend( + d.page_content for d in batch + ) ) try: ep_mod.embed_and_store_documents( @@ -337,7 +344,7 @@ class TestEmbedCheckpoint: attempt_id="sync-2", ) finally: - ep_mod.add_text_to_store_with_retry = original + ep_mod.add_texts_to_store_with_retry = original # All non-seed chunks re-embedded under the new attempt. assert captured == ["chunk-1", "chunk-2", "chunk-3", "chunk-4"] @@ -358,16 +365,18 @@ class TestEmbedCheckpoint: docs = _make_docs(4) captured: list[str] = [] - original = ep_mod.add_text_to_store_with_retry - ep_mod.add_text_to_store_with_retry = ( - lambda store, doc, sid: captured.append(doc.page_content) + original = ep_mod.add_texts_to_store_with_retry + ep_mod.add_texts_to_store_with_retry = ( + lambda store, batch, sid: captured.extend( + d.page_content for d in batch + ) ) try: ep_mod.embed_and_store_documents( docs, str(tmp_path), source_id, MagicMock(), ) finally: - ep_mod.add_text_to_store_with_retry = original + ep_mod.add_texts_to_store_with_retry = original # Resumed past last_index=1. assert captured == ["chunk-2", "chunk-3"] @@ -387,15 +396,15 @@ class TestEmbedCheckpoint: source_id = str(uuid.uuid4()) docs = _make_docs(1) - original = ep_mod.add_text_to_store_with_retry - ep_mod.add_text_to_store_with_retry = lambda store, doc, sid: None + original = ep_mod.add_texts_to_store_with_retry + ep_mod.add_texts_to_store_with_retry = lambda store, batch, sid: None try: ep_mod.embed_and_store_documents( docs, str(tmp_path), source_id, MagicMock(), attempt_id="att-single", ) finally: - ep_mod.add_text_to_store_with_retry = original + ep_mod.add_texts_to_store_with_retry = original row = pg_conn.execute( text( @@ -426,15 +435,15 @@ class TestEmbedCheckpoint: source_id = str(uuid.uuid4()) docs = _make_docs(4) - original = ep_mod.add_text_to_store_with_retry - ep_mod.add_text_to_store_with_retry = lambda store, doc, sid: None + original = ep_mod.add_texts_to_store_with_retry + ep_mod.add_texts_to_store_with_retry = lambda store, batch, sid: None try: ep_mod.embed_and_store_documents( docs, str(tmp_path), source_id, MagicMock(), attempt_id="att-multi", ) finally: - ep_mod.add_text_to_store_with_retry = original + ep_mod.add_texts_to_store_with_retry = original row = pg_conn.execute( text(