mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-03 07:11:56 +00:00
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:
1 parent
cb62dea701
commit
00be2c05ad
17 files changed
+449
-50
No files matched your search
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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
@@ -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")
|
||||
Reference in new issue
Block a user