From 00be2c05adfebf6668343ed70a1f9093a3c47d98 Mon Sep 17 00:00:00 2001 From: Alex Date: Fri, 28 Aug 2026 14:31:19 +0100 Subject: [PATCH] fix: make worker-delegated embedding survive the shipped deployments Query embedding moved to the Celery worker, but nothing that ships was updated to consume the queue it dispatches to. - Add `embeddings` to every worker `-Q` list (compose x3, k8s, devcontainer, sandbox README). Without it a search blocked for EMBEDDINGS_DELEGATE_TIMEOUT and then answered with no retrieved context, because classic_rag swallows the dispatch error and skips the source -- bad answers, not an error. - Skip the task_postrun heap reclaim for the embed task. The full gc.collect() was written for docling/torch parses; on a worker holding the ONNX model it measured ~86ms against ~8ms for the embed itself, a 9x slowdown of the round trip for a task that allocates a few kilobytes. - Resolve the installation pin in the re-embed script. It never imports application.app, so an install pinned in app_metadata with no EMBEDDINGS_NAME set -- every stock k8s deployment, whose manifests carry no embedding config -- would rewrite its whole index with the legacy default and stamp sources.model to match, then be told by the boot warning to run it again. - Fail fast for 30s after a failed dispatch. fanout.embed_questions falls back to letting each store embed its own query, so one dead-worker retrieval paid the timeout once in the fan-out and again per source. - Forget the task result. Nothing reads it back: the key is per-dispatch UUID, not content-addressed, so a repeated query mints another. Left alone every search leaked ~17KB for result_expires (7 days) into the Redis the broker shares -- on the bundled k8s manifest (1Gi, no maxmemory policy) that is an OOMKill that takes the broker with it. - Release the model ensure_vector_schema loads to read the width of an unregistered model, in a process that delegates and would never call it. The width still comes from the model, not the table, so the mismatch check the hook exists for keeps working. - Correct the docs that said otherwise: embeddings.md claimed the standard deployment worked unchanged, upgrading.mdx said no action was needed, and the settings table listed none of the three delegation settings. --- .devcontainer/devc-welcome.md | 2 +- application/celery_init.py | 19 ++- application/scripts/reembed.py | 11 ++ application/storage/db/bootstrap.py | 109 +++++++++++++----- .../vectorstore/embeddings_delegated.py | 68 ++++++++++- deployment/docker-compose-azure.yaml | 2 +- deployment/docker-compose-hub.yaml | 2 +- deployment/docker-compose.yaml | 11 +- .../k8s/deployments/docsgpt-deploy.yaml | 2 +- deployment/sandbox/README.md | 2 +- docs/content/Deploying/DocsGPT-Settings.mdx | 3 + docs/content/Models/embeddings.md | 4 +- docs/content/upgrading.mdx | 13 ++- tests/scripts/test_reembed.py | 31 +++++ .../db/test_bootstrap_vector_schema.py | 61 ++++++++++ tests/test_celery.py | 52 ++++++++- .../vectorstore/test_embeddings_delegated.py | 107 +++++++++++++++++ 17 files changed, 449 insertions(+), 50 deletions(-) diff --git a/.devcontainer/devc-welcome.md b/.devcontainer/devc-welcome.md index ad4103ba..aecfb84e 100644 --- a/.devcontainer/devc-welcome.md +++ b/.devcontainer/devc-welcome.md @@ -29,7 +29,7 @@ serves only the WSGI Flask app — it omits `/mcp` and the reconnect reader ### Celery (Task Queue) ```bash -celery -A application.app.celery worker -l INFO -Q docsgpt,parsing +celery -A application.app.celery worker -l INFO -Q docsgpt,parsing,embeddings ``` The `parsing` queue serves document parsing (the `read_document` tool / workflow diff --git a/application/celery_init.py b/application/celery_init.py index 6e400105..d3e33b1f 100644 --- a/application/celery_init.py +++ b/application/celery_init.py @@ -115,9 +115,24 @@ def _trim_native_heap() -> None: pass +# Tasks that allocate almost nothing and run on a latency-sensitive path, so +# the reclaim below costs far more than it recovers. Query embedding is one: +# measured at ~86 ms for the collect against ~8 ms for the embed itself on a +# worker holding the ONNX model, i.e. a 9x slowdown of the whole round trip. +_NO_RECLAIM_TASKS = frozenset({"application.vectorstore.embeddings_tasks.embed_texts"}) + + @task_postrun.connect -def _reclaim_memory_after_task(*args, **kwargs): - """Drop per-task allocations so the prefork child's RSS doesn't ratchet.""" +def _reclaim_memory_after_task(task=None, **kwargs): + """Drop per-task allocations so the prefork child's RSS doesn't ratchet. + + Skipped for the tasks in :data:`_NO_RECLAIM_TASKS`. This exists for the + large transient allocations docling/torch parsing makes; running a full + generational collect after a task that allocated a few kilobytes just + charges the next task for walking the whole heap. + """ + if getattr(task, "name", None) in _NO_RECLAIM_TASKS: + return gc.collect() torch = sys.modules.get("torch") if torch is not None: diff --git a/application/scripts/reembed.py b/application/scripts/reembed.py index d614c846..5f75f9e4 100644 --- a/application/scripts/reembed.py +++ b/application/scripts/reembed.py @@ -440,6 +440,17 @@ def main(argv: Optional[Sequence[str]] = None) -> int: args = build_parser().parse_args(argv) _log_setup(args.verbose) + # Which model this installation uses may live in ``app_metadata`` rather + # than the environment -- ``application.app`` resolves it at boot, and this + # script never imports that. Without this, an install pinned to granite + # with no EMBEDDINGS_NAME set (every stock Kubernetes deployment: the + # manifests carry no embedding config at all) would re-embed its whole + # index with the *legacy* code default and stamp ``sources.model`` to + # match -- the silent cross-model index this script exists to repair. + from application.storage.db.embeddings_pin import resolve_embeddings_pin + + resolve_embeddings_pin(logger) + # Embed in this process. ``EMBEDDINGS_DELEGATE_TO_WORKER`` exists to keep a # model out of the API, which serves one query at a time and holds the # model for nothing in between. This is the opposite case: a batch job that diff --git a/application/storage/db/bootstrap.py b/application/storage/db/bootstrap.py index 48dc6264..ca15dcd2 100644 --- a/application/storage/db/bootstrap.py +++ b/application/storage/db/bootstrap.py @@ -76,6 +76,43 @@ def ensure_database_ready( _run_migrations(log) +def _release_boot_only_embeddings(log: logging.Logger) -> None: + """Drop a model this hook loaded that the process will never use again. + + ``ensure_vector_schema`` has to run the model to learn the width of a model + the registry does not describe. ``EmbeddingsSingleton`` then caches it for + the life of the process -- correct when this process embeds, pure waste in + an API that delegates every embed to the worker, where it costs ~400 MB for + a small model and ~800 MB for a granite-sized one that is never called. + + Only the local-ONNX case is dropped. A ``RemoteEmbeddings`` is cheap and is + exactly what the process goes on to use, and with delegation off the model + would only be rebuilt on the first query. + + Bounds retention, not the transient peak: the load still happens, and the + ONNX Runtime arena may not return every page to the OS. + """ + from application.core.settings import settings + + if settings.EMBEDDINGS_BASE_URL: + return + if getattr(settings, "EMBEDDINGS_DELEGATE_TO_WORKER", False) is not True: + return + + import gc + + from application.vectorstore.base import EmbeddingsSingleton + + if EmbeddingsSingleton._instances.pop(settings.EMBEDDINGS_NAME, None) is None: + return + gc.collect() + log.info( + "ensure_vector_schema: released the embeddings model loaded to read its " + "width; this process delegates embedding to the worker and would never " + "have used it." + ) + + def ensure_vector_schema(*, logger: Optional[logging.Logger] = None) -> None: """Create the pgvector schema once at boot and verify its dimension. @@ -134,38 +171,6 @@ def ensure_vector_schema(*, logger: Optional[logging.Logger] = None) -> None: from application.vectorstore.model_registry import dimension_for dim: Optional[int] = dimension_for(settings.EMBEDDINGS_NAME) - if dim is None: - # An unregistered model only reports its width once something has run - # it. Build it in-process rather than through ``get_embeddings``: at - # boot there is no Celery task in flight, so a delegating client would - # dispatch to a worker that may not be up yet. - try: - from application.vectorstore.base import build_local_embeddings - - embedding = build_local_embeddings() - dim = getattr(embedding, "dimension", None) - if not dim: - # A remote client knows nothing about its server until it has - # called it, so ask once. Without this the table is sized at the - # default and the check below is skipped -- which is how a remote - # model of any other width silently got a vector(768) column, the - # exact failure this hook exists to catch. Milvus and Qdrant probe - # the same way. - dim = len(embedding.embed_query("dimension probe")) - except Exception as exc: # noqa: BLE001 — never block boot on the model - log.warning( - "ensure_vector_schema: could not determine the embedding width " - "(%s); creating the table with %d dimensions and skipping the " - "dimension check.", - exc, - DEFAULT_EMBEDDING_DIM, - ) - if dim is None: - log.warning( - "ensure_vector_schema: the embeddings model exposes no dimension; " - "using %d and skipping the dimension check.", - DEFAULT_EMBEDDING_DIM, - ) graph_enabled = bool(getattr(settings, "GRAPHRAG_ENABLED", False)) started = time.monotonic() @@ -184,6 +189,46 @@ def ensure_vector_schema(*, logger: Optional[logging.Logger] = None) -> None: finally: cursor.close() + if dim is None: + # An unregistered model only reports its width once something has + # run it. Build it in-process rather than through + # ``get_embeddings``: at boot there is no Celery task in flight, so + # a delegating client would dispatch to a worker that may not be up + # yet. The width must come from the model and not from the existing + # table -- reading the table would make the check below compare a + # value against itself, which is how a model of a different width + # silently inherits a table it does not fit. + try: + from application.vectorstore.base import build_local_embeddings + + embedding = build_local_embeddings() + dim = getattr(embedding, "dimension", None) + if not dim: + # A remote client knows nothing about its server until it + # has called it, so ask once. Without this the table is + # sized at the default and the check below is skipped -- + # which is how a remote model of any other width silently + # got a vector(768) column, the exact failure this hook + # exists to catch. Milvus and Qdrant probe the same way. + dim = len(embedding.embed_query("dimension probe")) + except Exception as exc: # noqa: BLE001 — never block boot on the model + log.warning( + "ensure_vector_schema: could not determine the embedding width " + "(%s); creating the table with %d dimensions and skipping the " + "dimension check.", + exc, + DEFAULT_EMBEDDING_DIM, + ) + finally: + _release_boot_only_embeddings(log) + + if dim is None: + log.warning( + "ensure_vector_schema: the embeddings model exposes no dimension; " + "using %d and skipping the dimension check.", + DEFAULT_EMBEDDING_DIM, + ) + PGVectorStore.create_schema(conn, dimension=dim or DEFAULT_EMBEDDING_DIM) if graph_enabled: from application.graphrag.store import ( diff --git a/application/vectorstore/embeddings_delegated.py b/application/vectorstore/embeddings_delegated.py index b952537a..2c72a61e 100644 --- a/application/vectorstore/embeddings_delegated.py +++ b/application/vectorstore/embeddings_delegated.py @@ -22,6 +22,7 @@ network hop rather than a broker round trip. from __future__ import annotations import logging +import time from typing import Any, List, Optional from application.core.settings import settings @@ -33,6 +34,41 @@ logger = logging.getLogger(__name__) #: it ``application.worker``, which pulls in the whole parsing stack. EMBED_TASK = "application.vectorstore.embeddings_tasks.embed_texts" +#: How long after a failed dispatch to fail fast instead of waiting out another +#: full ``EMBEDDINGS_DELEGATE_TIMEOUT``. Short enough that a worker restart is +#: picked up within one query, long enough to collapse the retries inside a +#: single retrieval into one timeout rather than one per source. +_FAILURE_COOLDOWN = 30.0 + +_NO_WORKER_HINT = ( + "Start a worker consuming it, point EMBEDDINGS_BASE_URL at an embedding " + "service, or set EMBEDDINGS_DELEGATE_TO_WORKER=false to load the model in " + "this process instead." +) + + +def _forget(result) -> None: + """Drop the task's stored vector from the result backend. + + Nothing ever reads it back. The key is ``celery-task-meta-``, minted + per dispatch rather than derived from the text, so a repeated query is a new + task and a new key -- the value is written once, read once by the ``get()`` + already waiting on it, then dead. Left alone it occupies ~17 KB for + ``result_expires`` (7 days), in the Redis the broker also runs on. + + Also releases the backend's pub/sub subscription for the task, which + ``get()`` alone does not. + + Never raises: the vector is already in hand, and a backend that cannot + delete must not fail the search. On the timeout path the worker may still + store its result afterwards, leaving one orphaned key -- no worse than not + forgetting at all, and bounded by the same expiry. + """ + try: + result.forget() + except Exception as exc: # noqa: BLE001 — cleanup must never fail a query + logger.debug("Could not forget the embed task result: %s", exc) + def _in_worker() -> bool: """True when a Celery task is executing in this process.""" @@ -52,6 +88,13 @@ class DelegatedEmbeddings: self.embeddings_key = embeddings_key self._local: Any = None self._dimension: Optional[int] = dimension_for(embeddings_name) + self._failed_at: Optional[float] = None + + def _cooldown_remaining(self) -> float: + """Seconds left of the fail-fast window after a failed dispatch.""" + if self._failed_at is None: + return 0.0 + return max(0.0, _FAILURE_COOLDOWN - (time.monotonic() - self._failed_at)) def _local_embeddings(self): """The in-process model, built once, for use inside a worker task.""" @@ -67,17 +110,34 @@ class DelegatedEmbeddings: queue = getattr(settings, "EMBEDDINGS_QUEUE", "embeddings") timeout = getattr(settings, "EMBEDDINGS_DELEGATE_TIMEOUT", 60) + + # A missing worker is a property of the deployment, not of this call, + # so once one dispatch has timed out the next is not worth another full + # timeout. Without this latch a single retrieval pays the timeout twice + # -- once in ``fanout.embed_questions``, then again per source when it + # falls back to letting each store embed its own query. + remaining = self._cooldown_remaining() + if remaining > 0: + raise RuntimeError( + f"Skipping the embed dispatch: a previous request to the {queue!r} " + f"queue failed and the {_FAILURE_COOLDOWN}s cooldown has " + f"{remaining:.0f}s left. {_NO_WORKER_HINT}" + ) + result = celery.send_task(EMBED_TASK, args=[texts, self.embeddings_name], queue=queue) try: - return result.get(timeout=timeout) + vectors = result.get(timeout=timeout) except Exception as exc: + self._failed_at = time.monotonic() raise RuntimeError( f"Embedding request to the Celery worker timed out or failed ({exc}). " f"A worker must be consuming the {queue!r} queue for retrieval to " - "work. Start one, point EMBEDDINGS_BASE_URL at an embedding " - "service, or set EMBEDDINGS_DELEGATE_TO_WORKER=false to load the " - "model in this process instead." + f"work. {_NO_WORKER_HINT}" ) from exc + finally: + _forget(result) + self._failed_at = None + return vectors def embed_documents(self, documents: List[str]) -> List[List[float]]: """Embed a list of texts, preserving order.""" diff --git a/deployment/docker-compose-azure.yaml b/deployment/docker-compose-azure.yaml index 3295a12c..4055646f 100644 --- a/deployment/docker-compose-azure.yaml +++ b/deployment/docker-compose-azure.yaml @@ -33,7 +33,7 @@ services: worker: build: ../application # `parsing` queue carries read_document/parse_document; required for its await to resolve. - command: celery -A application.app.celery worker -l INFO -Q docsgpt,parsing + command: celery -A application.app.celery worker -l INFO -Q docsgpt,parsing,embeddings env_file: - ../.env environment: diff --git a/deployment/docker-compose-hub.yaml b/deployment/docker-compose-hub.yaml index df41d038..fe691aee 100644 --- a/deployment/docker-compose-hub.yaml +++ b/deployment/docker-compose-hub.yaml @@ -40,7 +40,7 @@ services: user: root image: arc53/docsgpt:develop # `parsing` queue carries read_document/parse_document; required for its await to resolve. - command: celery -A application.app.celery worker -l INFO -B -Q docsgpt,parsing + command: celery -A application.app.celery worker -l INFO -B -Q docsgpt,parsing,embeddings env_file: - ../.env environment: diff --git a/deployment/docker-compose.yaml b/deployment/docker-compose.yaml index b027ac1a..97e0fe07 100644 --- a/deployment/docker-compose.yaml +++ b/deployment/docker-compose.yaml @@ -39,10 +39,13 @@ services: worker: user: root build: ../application - # Consumes the default queue AND the dedicated `parsing` queue (read_document / - # parse_document). Without `parsing` here the read_document await never resolves. - # For heavy/OCR parsing run a separate worker with `-Q parsing`. - command: celery -A application.app.celery worker -l INFO -B -Q docsgpt,parsing + # Consumes the default queue AND the dedicated `parsing` (read_document / + # parse_document) and `embeddings` (query embedding) queues. Without `parsing` + # the read_document await never resolves; without `embeddings` every search + # fails after EMBEDDINGS_DELEGATE_TIMEOUT, because EMBEDDINGS_DELEGATE_TO_WORKER + # is on by default. For heavy/OCR parsing run a separate worker with `-Q parsing`; + # to keep query latency off the ingest pool, another with `-Q embeddings`. + command: celery -A application.app.celery worker -l INFO -B -Q docsgpt,parsing,embeddings env_file: - ../.env environment: diff --git a/deployment/k8s/deployments/docsgpt-deploy.yaml b/deployment/k8s/deployments/docsgpt-deploy.yaml index fc7631fd..4e59d1a1 100644 --- a/deployment/k8s/deployments/docsgpt-deploy.yaml +++ b/deployment/k8s/deployments/docsgpt-deploy.yaml @@ -87,7 +87,7 @@ spec: image: arc53/docsgpt # `parsing` queue carries read_document/parse_document; required for its await to resolve. # For heavy/OCR parsing, run a separate deployment with `-Q parsing` (and GPU env). - command: ["celery", "-A", "application.app.celery", "worker", "-l", "INFO", "-n", "worker.%h", "-Q", "docsgpt,parsing"] + command: ["celery", "-A", "application.app.celery", "worker", "-l", "INFO", "-n", "worker.%h", "-Q", "docsgpt,parsing,embeddings"] resources: limits: memory: "4Gi" diff --git a/deployment/sandbox/README.md b/deployment/sandbox/README.md index 8de18717..af76f36a 100644 --- a/deployment/sandbox/README.md +++ b/deployment/sandbox/README.md @@ -207,7 +207,7 @@ libraries) so OCR-heavy parsing runs on a separate, optionally larger pool. worker must also consume `parsing`, or the tool's await never resolves: ```bash -celery -A application.app.celery worker -Q docsgpt,parsing -l INFO +celery -A application.app.celery worker -Q docsgpt,parsing,embeddings -l INFO ``` Tuning settings: `DOCUMENT_PARSE_TIMEOUT` (seconds the tool awaits before diff --git a/docs/content/Deploying/DocsGPT-Settings.mdx b/docs/content/Deploying/DocsGPT-Settings.mdx index 84b3bc6e..bb915ce8 100644 --- a/docs/content/Deploying/DocsGPT-Settings.mdx +++ b/docs/content/Deploying/DocsGPT-Settings.mdx @@ -447,6 +447,9 @@ See [Embeddings](/Models/embeddings) for full guidance. | `EMBEDDINGS_BASE_URL` | unset | Base URL of a remote OpenAI-compatible embeddings server. Setting it routes all embedding calls there. | | `EMBEDDINGS_KEY` | unset | Optional bearer token for the remote embeddings server. | | `EMBEDDINGS_MAX_INPUT_TOKENS` | unset | Truncate each remote embedding input to N tokens (guards servers that reject oversized inputs). | +| `EMBEDDINGS_DELEGATE_TO_WORKER` | `true` | Embed queries on the Celery worker instead of loading a model in the API. Requires a worker consuming `EMBEDDINGS_QUEUE`; set `false` to run the API standalone. Ignored when `EMBEDDINGS_BASE_URL` is set. | +| `EMBEDDINGS_QUEUE` | `embeddings` | Queue the query-embedding task is routed to. A worker started with an explicit `-Q` must list it. | +| `EMBEDDINGS_DELEGATE_TIMEOUT` | `60` | Seconds the API waits for the worker's vector before failing the search. | ## Tools Settings diff --git a/docs/content/Models/embeddings.md b/docs/content/Models/embeddings.md index ba15a5e1..ef070b6c 100644 --- a/docs/content/Models/embeddings.md +++ b/docs/content/Models/embeddings.md @@ -95,7 +95,9 @@ A local embedding model costs a few hundred megabytes of resident memory per pro `EMBEDDINGS_DELEGATE_TO_WORKER` (on by default) moves that work to the Celery worker: the API sends the text over the broker and gets the vector back, holding no model. Measured on a default install, the API process drops from ~657 MB to ~284 MB, and query embedding costs one broker round trip (~60 ms on a prefork worker). -Retrieval then depends on a worker consuming `EMBEDDINGS_QUEUE` (`embeddings` by default). A bare `celery worker` with no `-Q` consumes it along with everything else, so the standard deployment works unchanged — but its concurrency is shared with ingest, so a query can queue behind a long parse. Run a dedicated worker to isolate query latency: +Retrieval then depends on a worker consuming `EMBEDDINGS_QUEUE` (`embeddings` by default). A bare `celery worker` with no `-Q` consumes it along with everything else. **A worker started with an explicit `-Q` must list it** — the bundled Compose and Kubernetes manifests run `-Q docsgpt,parsing,embeddings` for exactly this reason. Omit it and every search blocks for `EMBEDDINGS_DELEGATE_TIMEOUT` and then returns an answer with no retrieved context, without raising. + +Sharing one worker also shares its concurrency with ingest, so a query can queue behind a long parse. Run a dedicated worker to isolate query latency: ```bash celery -A application.app.celery worker -Q embeddings diff --git a/docs/content/upgrading.mdx b/docs/content/upgrading.mdx index b6092cd7..25dd96c2 100644 --- a/docs/content/upgrading.mdx +++ b/docs/content/upgrading.mdx @@ -13,7 +13,18 @@ import { Callout } from 'nextra/components' ## Embedding models -DocsGPT now runs embeddings through [FastEmbed](https://github.com/qdrant/fastembed) (ONNX Runtime) instead of SentenceTransformer. The models are the same and the vectors are identical, so **an existing deployment needs no action** — `all-mpnet-base-v2` keeps working exactly as before. +DocsGPT now runs embeddings through [FastEmbed](https://github.com/qdrant/fastembed) (ONNX Runtime) instead of SentenceTransformer. The models are the same and the vectors are identical, so **your existing index needs no action** — `all-mpnet-base-v2` keeps working exactly as before. + + + **Your worker command does need one change.** Query embedding now runs on the Celery worker (`EMBEDDINGS_DELEGATE_TO_WORKER`, on by default), which keeps the API from loading a model of its own. If you start your worker with an explicit `-Q`, add the `embeddings` queue: + + ```diff + - celery -A application.app.celery worker -l INFO -Q docsgpt,parsing + + celery -A application.app.celery worker -l INFO -Q docsgpt,parsing,embeddings + ``` + + The bundled Compose and Kubernetes manifests already do this — pull them along with the code. Without it, every search blocks for `EMBEDDINGS_DELEGATE_TIMEOUT` (60s) and then answers with no retrieved context rather than raising, so the symptom is bad answers, not an error. To keep the model out of the worker too, set `EMBEDDINGS_BASE_URL`; to run the API on its own, set `EMBEDDINGS_DELEGATE_TO_WORKER=false`. + New installs default to `ibm-granite/granite-embedding-311m-multilingual-r2`: multilingual, a 32k-token context, and the same 768 dimensions. diff --git a/tests/scripts/test_reembed.py b/tests/scripts/test_reembed.py index 2b92378c..e6efec74 100644 --- a/tests/scripts/test_reembed.py +++ b/tests/scripts/test_reembed.py @@ -428,3 +428,34 @@ class TestRecordsTheModel: ), patch.object(reembed, "reembed_pgvector", side_effect=RuntimeError("boom")): assert reembed.run("pgvector", None, 64, False) == 1 record.assert_not_called() + + +class TestThePinIsResolved: + """The script must embed with the model the installation is pinned to. + + ``resolve_embeddings_pin`` runs in ``application.app``, which this script + never imports. An install pinned in ``app_metadata`` with no + ``EMBEDDINGS_NAME`` in the environment -- every stock Kubernetes + deployment, whose manifests carry no embedding config -- would otherwise + rewrite its whole index with the legacy code default and stamp + ``sources.model`` to match, creating the cross-model index this script + exists to repair. + """ + + def test_main_resolves_the_pin_before_reading_the_store(self): + order = [] + with patch( + "application.storage.db.embeddings_pin.resolve_embeddings_pin", + side_effect=lambda *a, **k: order.append("pin"), + ), patch.object(reembed.settings, "VECTOR_STORE", "pgvector", create=True), patch.object( + reembed, "run", side_effect=lambda *a, **k: (order.append("run"), 0)[1] + ): + assert reembed.main([]) == 0 + assert order == ["pin", "run"], "the pin must resolve before anything embeds" + + def test_an_unsupported_store_still_resolved_the_pin_first(self): + with patch( + "application.storage.db.embeddings_pin.resolve_embeddings_pin" + ) as pin, patch.object(reembed.settings, "VECTOR_STORE", "qdrant", create=True): + assert reembed.main([]) == 2 + pin.assert_called_once() diff --git a/tests/storage/db/test_bootstrap_vector_schema.py b/tests/storage/db/test_bootstrap_vector_schema.py index e9761b03..acea6e4b 100644 --- a/tests/storage/db/test_bootstrap_vector_schema.py +++ b/tests/storage/db/test_bootstrap_vector_schema.py @@ -326,3 +326,64 @@ class TestUnknownWidthIsProbed: vector_schema, _ = self._run(local, table_dimension=384) local.embed_query.assert_not_called() assert vector_schema.call_args.kwargs["dimension"] == 384 + + +@pytest.mark.unit +class TestBootLoadedModelIsReleased: + """The width probe must not leave a model resident in a delegating process. + + Reading ``.dimension`` off an unregistered model means loading it, and + ``EmbeddingsSingleton`` caches what it builds. In an API that delegates + every embed to the worker that cached copy is never called again — it is + several hundred megabytes held for the life of the process, which is the + cost ``EMBEDDINGS_DELEGATE_TO_WORKER`` exists to avoid. + """ + + def _run(self, vector_settings, *, delegate, base_url=None): + from application.vectorstore.base import EmbeddingsSingleton + + monkeyed = _embeddings(1024) + conn = MagicMock() + conn.cursor.return_value = MagicMock() + EmbeddingsSingleton._instances.pop("test-model", None) + + def _build(*_args, **_kwargs): + EmbeddingsSingleton._instances["test-model"] = monkeyed + return monkeyed + + with patch.object( + vector_settings, "EMBEDDINGS_DELEGATE_TO_WORKER", delegate + ), patch.object( + vector_settings, "EMBEDDINGS_BASE_URL", base_url + ), patch("psycopg.connect", return_value=conn), patch( + "application.vectorstore.model_registry.dimension_for", return_value=None + ), patch( + "application.vectorstore.base.build_local_embeddings", side_effect=_build + ), patch( + "application.vectorstore.pgvector.PGVectorStore.create_schema" + ) as vector_schema, patch( + "application.vectorstore.pgvector.PGVectorStore.table_dimension", + return_value=1024, + ): + ensure_vector_schema() + try: + return vector_schema, "test-model" in EmbeddingsSingleton._instances + finally: + EmbeddingsSingleton._instances.pop("test-model", None) + + def test_a_delegating_process_does_not_retain_it(self, vector_settings): + vector_schema, retained = self._run(vector_settings, delegate=True) + assert not retained, "a delegating API must not hold the model it probed" + assert vector_schema.call_args.kwargs["dimension"] == 1024 + + def test_a_process_that_embeds_locally_keeps_it(self, vector_settings): + _, retained = self._run(vector_settings, delegate=False) + assert retained, "without delegation the model is used, so evicting it "\ + "would only force a rebuild on the first query" + + def test_a_remote_client_is_kept(self, vector_settings): + _, retained = self._run( + vector_settings, delegate=True, base_url="http://embeddings:8080" + ) + assert retained, "a RemoteEmbeddings holds no model and is what the "\ + "process goes on to use" diff --git a/tests/test_celery.py b/tests/test_celery.py index 8994dbc4..afda82a2 100644 --- a/tests/test_celery.py +++ b/tests/test_celery.py @@ -1,4 +1,4 @@ -from unittest.mock import patch +from unittest.mock import MagicMock, patch import pytest from application.celery_init import make_celery @@ -222,3 +222,53 @@ def test_unparseable_file_raises_the_non_retryable_type(): with pytest.raises(DocumentParseError, match="No text could be extracted"): embed_and_store_documents([], "/tmp", "src", None) + + +@pytest.mark.unit +class TestReclaimIsSkippedForEmbeds: + """The post-task heap reclaim must not run on the query hot path. + + ``_reclaim_memory_after_task`` exists for the large transient allocations + docling/torch parsing makes. Query embedding became a Celery task, and a + full generational collect on a worker holding the ONNX model measured ~86 ms + against ~8 ms for the embed itself -- a 9x slowdown of the round trip for a + task that allocates a few kilobytes. + """ + + @staticmethod + def _collects(task_name): + from application.celery_init import _reclaim_memory_after_task + + task = MagicMock() + task.name = task_name + with patch("application.celery_init.gc.collect") as collect, patch( + "application.celery_init._trim_native_heap" + ): + _reclaim_memory_after_task(task=task, task_id="t", state="SUCCESS") + return collect.called + + def test_the_embed_task_is_skipped(self): + assert not self._collects("application.vectorstore.embeddings_tasks.embed_texts") + + def test_parsing_still_reclaims(self): + assert self._collects("application.api.user.tasks.parse_document") + + def test_ingest_still_reclaims(self): + assert self._collects("application.api.user.tasks.ingest") + + def test_an_unnamed_sender_still_reclaims(self): + """Unknown callers keep the old behaviour rather than silently skipping.""" + from application.celery_init import _reclaim_memory_after_task + + with patch("application.celery_init.gc.collect") as collect, patch( + "application.celery_init._trim_native_heap" + ): + _reclaim_memory_after_task(task_id="t", state="SUCCESS") + assert collect.called + + def test_the_skip_list_names_the_real_task(self): + """A renamed task must not silently start paying the collect again.""" + from application.celery_init import _NO_RECLAIM_TASKS + from application.vectorstore.embeddings_delegated import EMBED_TASK + + assert EMBED_TASK in _NO_RECLAIM_TASKS diff --git a/tests/vectorstore/test_embeddings_delegated.py b/tests/vectorstore/test_embeddings_delegated.py index b38b7310..8e91e25d 100644 --- a/tests/vectorstore/test_embeddings_delegated.py +++ b/tests/vectorstore/test_embeddings_delegated.py @@ -128,3 +128,110 @@ class TestGetEmbeddingsDispatch: with patch.object(base.settings, "EMBEDDINGS_BASE_URL", None): with patch.object(base.settings, "EMBEDDINGS_DELEGATE_TO_WORKER", True): assert base.get_embeddings("some/model") is base.get_embeddings("some/model") + + +class TestFailureCooldown: + """One dead-worker timeout per retrieval, not one per source. + + ``fanout.embed_questions`` swallows a dispatch failure and lets every store + embed its own query, so without a latch a single chat request pays + ``EMBEDDINGS_DELEGATE_TIMEOUT`` once in the fan-out and again per source. + A missing worker is a property of the deployment, not of the call. + """ + + @staticmethod + def _celery(side_effect): + result = MagicMock() + result.get.side_effect = side_effect + celery = MagicMock() + celery.send_task.return_value = result + return celery, result + + def test_only_the_first_call_waits_out_the_timeout(self, not_in_worker): + celery, _ = self._celery(TimeoutError("no worker")) + embeddings = DelegatedEmbeddings("granite-311m") + with patch.dict("sys.modules", {"application.celery_init": MagicMock(celery=celery)}): + for _ in range(4): + with pytest.raises(RuntimeError): + embeddings.embed_query("q") + assert celery.send_task.call_count == 1 + + def test_the_fast_failure_still_names_the_remedy(self, not_in_worker): + celery, _ = self._celery(TimeoutError("no worker")) + embeddings = DelegatedEmbeddings("granite-311m") + with patch.dict("sys.modules", {"application.celery_init": MagicMock(celery=celery)}): + with pytest.raises(RuntimeError): + embeddings.embed_query("q") + with pytest.raises(RuntimeError, match="EMBEDDINGS_DELEGATE_TO_WORKER=false"): + embeddings.embed_query("q") + + def test_the_latch_clears_once_the_worker_answers(self, not_in_worker): + celery, result = self._celery(TimeoutError("no worker")) + embeddings = DelegatedEmbeddings("granite-311m") + with patch.dict("sys.modules", {"application.celery_init": MagicMock(celery=celery)}): + with pytest.raises(RuntimeError): + embeddings.embed_query("q") + embeddings._failed_at = None # stand in for the cooldown elapsing + result.get.side_effect = None + result.get.return_value = [[0.5, 0.5]] + assert embeddings.embed_query("q") == [0.5, 0.5] + assert embeddings._cooldown_remaining() == 0.0 + + def test_a_healthy_worker_is_never_latched(self, not_in_worker): + celery, result = self._celery(None) + result.get.return_value = [[0.1, 0.2]] + embeddings = DelegatedEmbeddings("granite-311m") + with patch.dict("sys.modules", {"application.celery_init": MagicMock(celery=celery)}): + for _ in range(3): + assert embeddings.embed_query("q") == [0.1, 0.2] + assert celery.send_task.call_count == 3 + + +class TestTheResultIsForgotten: + """A query vector must not outlive the query that asked for it. + + ``result_expires`` is 7 days and ``embed_texts`` stores its result, but the + key is ``celery-task-meta-`` -- minted per dispatch, never derived + from the text -- so nothing reads it back and a repeated query mints + another. Without ``forget()`` every search leaks ~17 KB into the Redis the + broker shares for a week. + """ + + @staticmethod + def _celery(side_effect=None, value=None): + result = MagicMock() + result.get.side_effect = side_effect + result.get.return_value = value + celery = MagicMock() + celery.send_task.return_value = result + return celery, result + + def test_a_successful_embed_forgets_its_result(self, not_in_worker): + celery, result = self._celery(value=[[0.1, 0.2]]) + embeddings = DelegatedEmbeddings("granite-311m") + with patch.dict("sys.modules", {"application.celery_init": MagicMock(celery=celery)}): + assert embeddings.embed_query("q") == [0.1, 0.2] + result.forget.assert_called_once() + + def test_a_failed_embed_still_forgets(self, not_in_worker): + celery, result = self._celery(side_effect=TimeoutError("no worker")) + embeddings = DelegatedEmbeddings("granite-311m") + with patch.dict("sys.modules", {"application.celery_init": MagicMock(celery=celery)}): + with pytest.raises(RuntimeError): + embeddings.embed_query("q") + result.forget.assert_called_once() + + def test_a_backend_that_cannot_delete_does_not_fail_the_query(self, not_in_worker): + celery, result = self._celery(value=[[0.3, 0.4]]) + result.forget.side_effect = ConnectionError("backend down") + embeddings = DelegatedEmbeddings("granite-311m") + with patch.dict("sys.modules", {"application.celery_init": MagicMock(celery=celery)}): + assert embeddings.embed_query("q") == [0.3, 0.4] + + def test_forgetting_does_not_mask_the_dispatch_failure(self, not_in_worker): + celery, result = self._celery(side_effect=TimeoutError("no worker")) + result.forget.side_effect = ConnectionError("backend down") + embeddings = DelegatedEmbeddings("granite-311m") + with patch.dict("sys.modules", {"application.celery_init": MagicMock(celery=celery)}): + with pytest.raises(RuntimeError, match="timed out or failed"): + embeddings.embed_query("q")