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.
This commit is contained in:
Alex committed 2026-08-28 14:31:19 +01:00
1 parent cb62dea701
commit 00be2c05ad
17 files changed
+449 -50

No files matched your search

+1 -1
View File
@@ -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
+17 -2
View File
@@ -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:
+11
View File
@@ -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
+77 -32
View File
@@ -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 (
@@ -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-<uuid>``, 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."""
+1 -1
View File
@@ -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:
+1 -1
View File
@@ -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:
+7 -4
View File
@@ -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:
@@ -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"
+1 -1
View File
@@ -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
@@ -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
+3 -1
View File
@@ -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
+12 -1
View File
@@ -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.
<Callout type="warning">
**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`.
</Callout>
New installs default to `ibm-granite/granite-embedding-311m-multilingual-r2`: multilingual, a 32k-token context, and the same 768 dimensions.
+31
View File
@@ -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()
@@ -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"
+51 -1
View File
@@ -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
@@ -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-<uuid>`` -- 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")