From 47e53a71c4e064db54702672e970b17e2a4a1d42 Mon Sep 17 00:00:00 2001 From: Alex Date: Fri, 28 Aug 2026 12:23:18 +0100 Subject: [PATCH] feat: embed on the worker, and stop batching the ONNX pass The API embeds every query it serves, so it held its own copy of the model: ~890 MB it never needed. EMBEDDINGS_DELEGATE_TO_WORKER (on by default) sends the text to the Celery worker instead and gets the vector back, taking an API process from 1176 MB to 285 MB with no ONNX Runtime imported at all. The client embeds locally when it finds itself inside a worker task, so the worker never dispatches to itself -- the same self-deadlock DOCUMENT_PARSE_QUEUE avoids on the parsing side. EMBEDDINGS_BASE_URL still wins over it, and remains the right answer for production. ensure_vector_schema was constructing the embeddings instance purely to read .dimension off it, loading several hundred MB of ONNX into every API and worker process at import. For a model the registry describes that is a lookup; only an unregistered name now falls back to loading. EMBEDDINGS_BATCH_SIZE was sizing two unrelated things: chunks per store transaction (and per remote embed request) and documents per ONNX forward pass. Each pass pads every input up to its longest, and that waste grows with the square of chunk length, so at the 1250-token default a batch of 32 peaked at 6.6 GB and took 326s where a batch of 1 peaked at 2.9 GB and took 90s. The forward pass is now sized by EMBEDDINGS_MODEL_BATCH_SIZE, defaulting to 1; storage and remote batching are unchanged at 32. reembed embeds in-process: a batch job that walks the whole index should not round-trip every chunk through a broker, and loading the model there reports a real failure instead of timing out against an empty queue. Also drops the mpnet zip download from the docs and the devcontainer, which pointed at a SentenceTransformers export with no ONNX graph and had been inert since the FastEmbed swap; corrects the claim that any sentence-transformers model works; and settles the Configuring/Settings pages on what the registry and the repository metadata actually decide. --- .devcontainer/post-create-command.sh | 9 +- .env-template | 17 + AGENTS.md | 24 +- application/celeryconfig.py | 15 +- application/core/settings.py | 303 +++++++----------- application/scripts/reembed.py | 11 + application/storage/db/bootstrap.py | 37 ++- application/vectorstore/base.py | 43 +++ .../vectorstore/embeddings_delegated.py | 118 +++++++ application/vectorstore/embeddings_local.py | 19 +- application/vectorstore/embeddings_tasks.py | 29 ++ .../Deploying/Development-Environment.mdx | 12 +- docs/content/Deploying/DocsGPT-Settings.mdx | 9 +- docs/content/Models/embeddings.md | 69 ++-- tests/conftest.py | 15 + tests/scripts/test_reembed.py | 16 + .../db/test_bootstrap_vector_schema.py | 47 ++- .../vectorstore/test_embeddings_delegated.py | 130 ++++++++ tests/vectorstore/test_embeddings_local.py | 8 +- .../vectorstore/test_pgvector_live_schema.py | 14 + 20 files changed, 679 insertions(+), 266 deletions(-) create mode 100644 application/vectorstore/embeddings_delegated.py create mode 100644 application/vectorstore/embeddings_tasks.py create mode 100644 tests/vectorstore/test_embeddings_delegated.py diff --git a/.devcontainer/post-create-command.sh b/.devcontainer/post-create-command.sh index 597b985e..ec9e5fb9 100755 --- a/.devcontainer/post-create-command.sh +++ b/.devcontainer/post-create-command.sh @@ -21,12 +21,9 @@ else fi -mkdir -p model -if [ ! -d model/all-mpnet-base-v2 ]; then - wget -q https://d3dg1063dc54p9.cloudfront.net/models/embeddings/mpnet-base-v2.zip -O model/mpnet-base-v2.zip - unzip -q model/mpnet-base-v2.zip -d model - rm model/mpnet-base-v2.zip -fi +# The embedding model is fetched on first use and cached, so nothing to download +# here. For an offline container, run `python -m application.scripts.prefetch_models` +# after the install below. pip install -r application/requirements.txt cd frontend npm install --include=dev \ No newline at end of file diff --git a/.env-template b/.env-template index 52d0e717..c62cf3c2 100644 --- a/.env-template +++ b/.env-template @@ -28,6 +28,23 @@ EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2 EMBEDDINGS_BASE_URL= EMBEDDINGS_KEY= +# Run the embedding model on the Celery worker instead of in every process that +# embeds. The API embeds each query it serves, so without this it holds its own +# copy of the model (~370 MB more resident). Costs a broker round trip per +# query. Retrieval then needs a worker consuming EMBEDDINGS_QUEUE -- set this to +# false if you run the API on its own. +# EMBEDDINGS_DELEGATE_TO_WORKER=true +# EMBEDDINGS_QUEUE=embeddings +# EMBEDDINGS_DELEGATE_TIMEOUT=60 + +# Documents per local ONNX forward pass. Each pass pads every input up to the +# longest one in it, and that waste grows with the square of chunk length, so +# larger is not faster here: at the 1250-token default chunk size, 32 peaked at +# 6.6 GB and took 326s, while 1 peaked at 2.9 GB and took 90s. Raise it only if +# your chunks are short and uniform. Distinct from EMBEDDINGS_BATCH_SIZE, which +# is chunks per store transaction / per remote embed request. +# EMBEDDINGS_MODEL_BATCH_SIZE=1 + #For Azure (you can delete it if you don't use Azure) OPENAI_API_BASE= OPENAI_API_VERSION= diff --git a/AGENTS.md b/AGENTS.md index b3fe6c4e..052b0c87 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -54,23 +54,37 @@ Production uses `gunicorn -k uvicorn_worker.UvicornWorker` against the same `application.asgi:asgi_app` target; see `application/Dockerfile` for the full flag set. -Run the Celery worker in a separate terminal (if needed): +Run the Celery worker in a separate terminal: ```bash celery -A application.app.celery worker -l INFO ``` +**The worker is required for retrieval, not optional.** `EMBEDDINGS_DELEGATE_TO_WORKER` +defaults on, so the API embeds each query by dispatching to the worker rather than +loading a model of its own — which keeps the API process around 285 MB instead of +1.2 GB. Without a worker consuming `EMBEDDINGS_QUEUE`, every search fails after +`EMBEDDINGS_DELEGATE_TIMEOUT`. To run the API on its own, either set +`EMBEDDINGS_DELEGATE_TO_WORKER=false` (loads the model in-process) or point +`EMBEDDINGS_BASE_URL` at an embeddings service. + On macOS, prefer the solo pool for Celery: ```bash python -m celery -A application.app.celery worker -l INFO --pool=solo ``` +Note that `--pool=solo` costs roughly 350 ms per query embed against ~55 ms on the +default prefork pool — nearly all of it the solo worker picking the message up, not +the embedding itself. That only affects local dev; production runs prefork. + A bare worker (no `-Q`) consumes every configured queue, so one worker does the -whole job — app tasks and document parsing (the `read_document` tool / workflow -native-file parse) alike. Use `-Q` only to split load: run the main worker with -`-Q docsgpt` and a dedicated (e.g. GPU-enabled) parser worker with `-Q parsing` -for heavy OCR. +whole job — app tasks, query embedding, and document parsing (the `read_document` +tool / workflow native-file parse) alike. Use `-Q` only to split load: run the main +worker with `-Q docsgpt`, a dedicated (e.g. GPU-enabled) parser worker with +`-Q parsing` for heavy OCR, and `-Q embeddings` to keep query latency off the ingest +pool. Note the main `ingest` task parses in-process on `docsgpt`; only +`read_document` is routed to `parsing`. ### Frontend diff --git a/application/celeryconfig.py b/application/celeryconfig.py index fd696402..98584081 100644 --- a/application/celeryconfig.py +++ b/application/celeryconfig.py @@ -13,7 +13,10 @@ result_serializer = 'json' accept_content = ['json'] # Autodiscover tasks -imports = ('application.api.user.tasks',) +imports = ( + 'application.api.user.tasks', + 'application.vectorstore.embeddings_tasks', +) # Project-scoped queue so a stray sibling worker on the same broker # (other repo, same default ``celery`` queue) can't grab DocsGPT tasks. @@ -25,8 +28,13 @@ task_default_routing_key = "docsgpt" # Celery worker (headless/scheduled agent) is served by a separate parsing worker # and never self-deadlocks the awaiting worker. The tool also passes the queue at # apply_async time, so this routing is the default for any other enqueuer. +# Query embedding gets its own queue for the same reason parsing does: a query +# waiting behind a multi-minute ingest is a query that has timed out. A bare +# worker still consumes it, but its concurrency is shared -- run a separate +# ``-Q embeddings`` worker to actually isolate query latency from ingest. task_routes = { "application.api.user.tasks.parse_document": {"queue": settings.DOCUMENT_PARSE_QUEUE}, + "application.vectorstore.embeddings_tasks.embed_texts": {"queue": settings.EMBEDDINGS_QUEUE}, } # Declare every queue so a bare ``celery worker`` (no -Q) consumes ALL of them — @@ -34,7 +42,10 @@ task_routes = { # heavy OCR isolated run one worker with ``-Q docsgpt`` and another with # ``-Q parsing``. (dict.fromkeys dedupes if DOCUMENT_PARSE_QUEUE == "docsgpt".) task_queues = tuple( - Queue(name) for name in dict.fromkeys(["docsgpt", settings.DOCUMENT_PARSE_QUEUE]) + Queue(name) + for name in dict.fromkeys( + ["docsgpt", settings.DOCUMENT_PARSE_QUEUE, settings.EMBEDDINGS_QUEUE] + ) ) beat_scheduler = "redbeat.RedBeatScheduler" diff --git a/application/core/settings.py b/application/core/settings.py index deb2e3ab..d88e6bcc 100644 --- a/application/core/settings.py +++ b/application/core/settings.py @@ -33,10 +33,8 @@ class Settings(BaseSettings): OIDC_GROUPS_CLAIM: str = "groups" # ID-token/userinfo claim carrying group membership OIDC_ADMIN_GROUPS: Optional[str] = None # comma-separated groups granted admin; unset = no OIDC admin mapping - # RBAC (admin/user roles). Persisted admin grants live in the user_roles - # table and apply only under AUTH_TYPE=oidc. LOCAL_MODE_ADMIN is the only - # non-DB admin path and applies only to AUTH_TYPE=None (no-auth self-host). - # It MUST stay False in any networked deployment. + # RBAC: persisted admin grants live in user_roles (AUTH_TYPE=oidc only). This is the + # only non-DB admin path, for AUTH_TYPE=None self-host. MUST stay False if networked. LOCAL_MODE_ADMIN: bool = False # SCIM 2.0 provisioning (IdP-driven user create/deactivate at /scim/v2) @@ -45,59 +43,56 @@ class Settings(BaseSettings): LLM_PROVIDER: str = "docsgpt" LLM_NAME: Optional[str] = None # if LLM_PROVIDER is openai, LLM_NAME can be gpt-4 or gpt-3.5-turbo - # Deliberately the legacy model, not the current recommendation. An - # existing deployment that never pinned EMBEDDINGS_NAME falls through to - # this default, and its stored vectors were produced by this model -- - # changing the default here would silently retrieve against a different - # vector space (both are 768-dimensional, so no dimension check fires). - # New installs get granite from .env-template / setup.sh; existing ones - # switch by setting this and running application.scripts.reembed. + # Legacy model on purpose: an install that never pinned this has vectors from it, and + # granite is the same width so a swap would fail silently. New installs get granite from + # .env-template; existing ones switch by setting this and running application.scripts.reembed. EMBEDDINGS_NAME: str = "huggingface_sentence-transformers/all-mpnet-base-v2" EMBEDDINGS_BASE_URL: Optional[str] = None # Remote embeddings API URL (OpenAI-compatible) EMBEDDINGS_KEY: Optional[str] = None # api key for embeddings (if using openai, just copy API_KEY) EMBEDDINGS_MAX_INPUT_TOKENS: Optional[int] = None # truncate each remote embed input to N tokens (overflow lost) - EMBEDDINGS_BATCH_SIZE: int = 32 # chunks per embed request during ingest (1 = legacy per-chunk behaviour) - # Intra-op threads for the local ONNX runner. None = every core. ONNX Runtime - # scales sub-linearly across threads, so several single-threaded worker - # processes beat one many-threaded process on the same cores. + EMBEDDINGS_BATCH_SIZE: int = 32 # chunks per store transaction / remote embed request + # Documents per local ONNX forward pass. Each pass pads to its longest input, and that + # waste grows with the square of chunk length: at 1250 tokens, 32 peaked at 6.6 GB, 1 at 2.9 GB. + EMBEDDINGS_MODEL_BATCH_SIZE: int = 1 + # Intra-op threads for the local ONNX runner; None = every core. It scales sub-linearly, + # so several single-threaded workers beat one many-threaded process on the same cores. EMBEDDINGS_THREADS: Optional[int] = None EMBEDDINGS_CACHE_DIR: Optional[str] = None # where FastEmbed caches model artifacts - # How a local model turns token vectors into one vector, and whether the - # result is L2-normalised. Both are normally read from the model's own - # repository; set these only for a repository that declares neither, or to - # override what it declares. "cls" or "mean". + # Pooling ("cls"/"mean") and L2 normalisation. Read from the model's own repository; + # set these only for a repository that declares neither, or to override what it declares. EMBEDDINGS_POOLING: Optional[str] = None EMBEDDINGS_NORMALIZE: Optional[bool] = None + # Embed on the worker so the API holds no model (~890 MB), at one broker round trip per + # query. Ignored when EMBEDDINGS_BASE_URL is set, which is the better answer for production. + EMBEDDINGS_DELEGATE_TO_WORKER: bool = True + EMBEDDINGS_QUEUE: str = "embeddings" # queue the embed task is routed to + EMBEDDINGS_DELEGATE_TIMEOUT: int = 60 # seconds to wait for the worker GITHUB_INGEST_MAX_FILE_BYTES: int = 1048576 # skip repo blobs larger than this (0 = no cap) GITHUB_INGEST_MAX_WORKERS: int = 8 # parallel file fetches per GitHub repo ingest - # Optional directory of operator-supplied model YAMLs, loaded after the - # built-in catalog under application/core/models/. Later wins on + # Operator-supplied model YAMLs, loaded after the built-in catalog; later wins on # duplicate model id. See application/core/models/README.md. MODELS_CONFIG_DIR: Optional[str] = None CELERY_BROKER_URL: str = "redis://localhost:6379/0" CELERY_RESULT_BACKEND: str = "redis://localhost:6379/1" - # Prefetch=1 caps SIGKILL loss to one task. Visibility timeout must exceed - # the longest legitimate task runtime (ingest, agent webhook) but stay - # short enough that SIGKILLed tasks redeliver promptly. 1h matches Onyx - # and Dify defaults; long ingests can override via env. + # Prefetch=1 caps SIGKILL loss to one task. Visibility timeout must exceed the longest + # legitimate task runtime but stay short enough that SIGKILLed tasks redeliver promptly. CELERY_WORKER_PREFETCH_MULTIPLIER: int = 1 CELERY_VISIBILITY_TIMEOUT: int = 3600 - # Recycle the prefork worker child once its resident size crosses this many - # kilobytes — backstops native-heap growth from docling/torch parsing. 0 disables. + # Recycle a prefork child past this resident size in KB; backstops docling/torch heap growth. + # Checked between tasks, so it does not bound the peak within one. 0 disables. CELERY_WORKER_MAX_MEMORY_PER_CHILD: int = 4194304 - # Recycle the child after this many tasks; 0 disables (memory cap is the primary knob). - CELERY_WORKER_MAX_TASKS_PER_CHILD: int = 0 + CELERY_WORKER_MAX_TASKS_PER_CHILD: int = 0 # recycle after N tasks; 0 disables # Only consulted when VECTOR_STORE=mongodb or when running scripts/db/backfill.py; user data lives in Postgres. MONGO_URI: Optional[str] = None # User-data Postgres DB. POSTGRES_URI: Optional[str] = None - # On app startup, apply pending Alembic migrations. Default ON for dev; disable in prod if you manage schema out-of-band. + # On startup, apply pending Alembic migrations. Disable if you manage schema out-of-band. AUTO_MIGRATE: bool = True - # On app startup, create the target Postgres database if it's missing (requires CREATEDB privilege). Dev-friendly default. + # On startup, create the target Postgres database if missing (needs CREATEDB privilege). AUTO_CREATE_DB: bool = True - # On app startup, create the pgvector/graph tables and verify the embedding dimension. Set False to manage them - # out-of-band — there is no Alembic migration for the vector DB because it may be a separate cluster. + # On startup, create the pgvector/graph tables and verify the embedding dimension. No Alembic + # migration covers the vector DB (it may be a separate cluster); set False to manage it yourself. AUTO_VECTOR_SCHEMA: bool = True LLM_PATH: str = os.path.join(current_dir, "models/docsgpt-7b-f16.gguf") DEFAULT_MAX_HISTORY: int = 150 @@ -112,8 +107,7 @@ class Settings(BaseSettings): "request_limit": 500, } UPLOAD_FOLDER: str = "inputs" - # Public upload request cap is applied by Flask before multipart parsing. - # The per-file cap is also enforced while copying each controlled stream. + # Request cap is applied by Flask before multipart parsing; the per-file cap also while copying. UPLOAD_MAX_REQUEST_BYTES: int = Field(default=256 * 1024 * 1024, gt=0) UPLOAD_MAX_FILE_BYTES: int = Field(default=100 * 1024 * 1024, gt=0) PARSE_SPEC_MAX_BYTES: int = Field(default=10 * 1024 * 1024, gt=0) @@ -126,49 +120,35 @@ class Settings(BaseSettings): PARSE_IMAGE_REMOTE: bool = False DOCLING_OCR_ENABLED: bool = False # Enable OCR for docling parsers (PDF, images) DOCLING_OCR_ATTACHMENTS_ENABLED: bool = False # Enable OCR for docling when parsing attachments - # Pages docling's threaded pipeline buffers in flight; the library - # default (100) drives worker RSS to ~3 GB on a mid-size PDF. + # Pages docling buffers in flight; its default of 100 drives worker RSS to ~3 GB on a mid-size PDF. DOCLING_PIPELINE_QUEUE_MAX_SIZE: int = 2 DOCLING_COMPILE_TORCH_MODELS: bool = False DOCLING_TABULAR_MAX_BYTES: int = 2_000_000 DOCLING_MARKUP_MAX_BYTES: int = 8_000_000 - # Chars-per-page floor below which an OCR'd PDF/image parse is treated as a docling - # pipeline dropout (long-running workers were observed returning zero characters for - # every scanned page after a long scanned PDF, with no error) rather than as content. - # Such a parse is retried once on a fresh full-page-OCR converter and then fails - # loudly instead of indexing an empty document. 0 disables the guard. + # Chars-per-page floor below which an OCR'd parse is treated as a docling dropout rather than + # content: retried once on a fresh converter, then failed loudly instead of indexing an empty + # document. Long-running workers were seen returning zero chars per page with no error. 0 disables. DOCLING_OCR_MIN_CHARS_PER_PAGE: int = 20 - # Read PDF *attachments* via their embedded text layer (pypdfium2) instead - # of docling, falling back to docling when there is no text layer to read. - # Attachments go into a prompt, so docling's structural markdown earns far - # less than the tens of seconds per file it costs; source ingestion is - # unaffected and always uses docling, because chunking and retrieval do - # depend on that structure. + # Read PDF *attachments* via their text layer (pypdfium2), falling back to docling when there + # is none. Attachments go into a prompt, where docling's structure does not earn its tens of + # seconds per file. Source ingestion is unaffected and always uses docling. ATTACHMENT_PDF_TEXT_FAST_PATH: bool = True - # Median characters per sampled page below which a PDF is treated as a scan - # and handed to docling. Measured separation on real uploads: scans at - # 0-17 chars/page, text-layer documents at 433-6834. + # Median chars per sampled page below which a PDF is treated as a scan and handed to docling. + # Measured on real uploads: scans at 0-17 chars/page, text-layer documents at 433-6834. ATTACHMENT_PDF_TEXT_MIN_MEDIAN_CHARS: int = 32 ATTACHMENT_TEXT_MAX_BYTES: int = 5_000_000 AGENT_IMAGE_MAX_BYTES: int = 5_000_000 AGENT_IMAGE_MAX_PIXELS: int = 16_777_216 VECTOR_STORE: str = "faiss" # "faiss" or "elasticsearch" or "qdrant" or "milvus" or "lancedb" or "pgvector" - # Allow-list of retriever keys an agent may use. Values must match the - # ``RetrieverCreator.retrievers`` registry keys (``classic`` / ``default``), + # Retriever keys an agent may use; must match RetrieverCreator.retrievers registry keys, # NOT the legacy ``classic_rag`` label which never matched the registry. RETRIEVERS_ENABLED: list = ["classic", "default"] - # Concurrent per-source vector searches within one retrieval (multi-source chats); - # the query is embedded once and shared across sources. + # Concurrent per-source searches in one retrieval; the query is embedded once and shared. RETRIEVAL_MAX_PARALLEL_SOURCES: int = 4 - # Kill-switch for per-source retrieval dispatch. When False the retrieval - # path collapses to today's single-retriever behavior (consumed by the - # Dispatcher in a later change; defined here so the flag exists up front). + # Kill-switch for per-source retrieval dispatch; False collapses to a single retriever. PER_SOURCE_RETRIEVAL_ENABLED: bool = True - # Flagship GraphRAG flag. Reserved and unused for now; gates graph-aware - # ingestion/retrieval when that feature lands. - GRAPHRAG_ENABLED: bool = False - # Model for ingest-time graph extraction; None reuses the instance default - # model (LLM_PROVIDER/LLM_NAME). Operator-overridable (e.g. a cheaper model). + GRAPHRAG_ENABLED: bool = False # gates graph-aware ingestion/retrieval + # Model for ingest-time graph extraction; None reuses LLM_PROVIDER/LLM_NAME. GRAPHRAG_EXTRACTION_MODEL: Optional[str] = None # Hard cap on chunks extracted per source (cost control). GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION: int = 2000 @@ -231,9 +211,8 @@ class Settings(BaseSettings): ELASTIC_URL: Optional[str] = None # url for elasticsearch ELASTIC_INDEX: Optional[str] = "docsgpt" # index name for elasticsearch - # Legacy AWS credentials from the retired SageMaker LLM provider. - # Still read as a deprecated fallback by S3 storage (see the S3_* - # block below); do not use for new deployments. + # Legacy AWS credentials from the retired SageMaker provider. Still read as a deprecated + # fallback by S3 storage; do not use for new deployments. SAGEMAKER_REGION: Optional[str] = None SAGEMAKER_ACCESS_KEY: Optional[str] = None SAGEMAKER_SECRET_KEY: Optional[str] = None @@ -253,16 +232,11 @@ class Settings(BaseSettings): QDRANT_PATH: Optional[str] = None QDRANT_DISTANCE_FUNC: str = "Cosine" - # PGVector vectorstore config. Write the URI in whichever form you - # prefer — ``postgres://``, ``postgresql://``, or even the SQLAlchemy - # dialect form (``postgresql+psycopg://``) are all accepted and - # normalized internally for ``psycopg.connect()``. + # PGVector config. postgres://, postgresql:// and postgresql+psycopg:// are all accepted + # and normalized internally for psycopg.connect(). PGVECTOR_CONNECTION_STRING: Optional[str] = None - # Per-process psycopg connection pool for the vector store; 0 = one direct - # connection per store instance, the legacy behaviour. - PGVECTOR_POOL_MAX_SIZE: int = 8 - # IVFFlat probes for vector search. ``None`` derives sqrt(lists) from the - # index itself; set an integer to pin it. Higher = better recall, more scan. + PGVECTOR_POOL_MAX_SIZE: int = 8 # per-process pool; 0 = one direct connection per store + # IVFFlat probes; None derives sqrt(lists) from the index. Higher = better recall, more scan. PGVECTOR_IVFFLAT_PROBES: Optional[int] = None # Milvus vectorstore config MILVUS_COLLECTION_NAME: Optional[str] = "docsgpt" @@ -276,11 +250,8 @@ class Settings(BaseSettings): FLASK_DEBUG_MODE: bool = False STORAGE_TYPE: str = "local" # local or s3 - # S3-compatible object storage (used when STORAGE_TYPE=s3). Works with AWS - # S3 and any S3-compatible service (MinIO, Cloudflare R2, Backblaze B2, - # DigitalOcean Spaces, ...). For non-AWS services, set S3_ENDPOINT_URL and - # usually S3_PATH_STYLE=true. The SAGEMAKER_* credentials are still read as - # a deprecated fallback for backward compatibility. + # S3-compatible object storage (STORAGE_TYPE=s3): AWS S3, MinIO, R2, B2, Spaces, ... + # For non-AWS, set S3_ENDPOINT_URL and usually S3_PATH_STYLE=true. S3_BUCKET_NAME: str = "docsgpt-test-bucket" S3_ENDPOINT_URL: Optional[str] = None # custom endpoint for S3-compatible services; omit for AWS S3_ACCESS_KEY_ID: Optional[str] = None @@ -309,34 +280,21 @@ class Settings(BaseSettings): # Tool pre-fetch settings ENABLE_TOOL_PREFETCH: bool = True - # When True, OpenAI Responses API calls are persisted server-side - # (store=true) so a previous_response_id can chain turns. When False - # (the default) Responses calls are stateless (store=false) and any - # reasoning is carried across the in-turn tool loop via encrypted - # reasoning items instead. + # True persists Responses API calls server-side so previous_response_id can chain turns. + # False keeps them stateless, carrying reasoning across the tool loop as encrypted items. OPENAI_RESPONSES_STORE: bool = False OPENAI_REASONING_SUMMARY: str = "auto" - # OpenAI-compatible clients can identify a logical chat with session - # headers even though chat-completions itself has no conversation field. + # Lets OpenAI-compatible clients identify a logical chat by session header, which + # chat-completions itself has no field for. V1_SESSION_TTL_SECONDS: int = 24 * 60 * 60 - - # Optional cheaper model for first-party conversation titles. When unset, - # listed conversations use their answer model, but title work is still - # dispatched off the response path. + # Optional cheaper model for conversation titles; unset reuses the answer model. TITLE_MODEL_ID: Optional[str] = None - # Config-free tools on by default in agentless chats. ``scheduler`` is - # dual-registered (also in ``BUILTIN_AGENT_TOOLS``) so the same synthetic id - # resolves whether reached via defaults or the agent picker. - # - # ``code_executor`` and ``artifact_generator`` belong here too on any - # deployment that runs a sandbox (see SANDBOX_BACKEND / SANDBOX_GATEWAY_URL): - # ``artifact_generator`` renders the .docx/.pdf/.xlsx/.pptx files users ask - # chat for. They are left out of the shipped default because both execute - # through the sandbox runner and would fail on every call without one — add - # them explicitly once a runner is configured: - # DEFAULT_CHAT_TOOLS = [..., "code_executor", "artifact_generator"] + # Config-free tools on by default in agentless chats. ``scheduler`` is dual-registered in + # BUILTIN_AGENT_TOOLS so one synthetic id resolves via defaults or the agent picker. + # Add "code_executor" and "artifact_generator" once a sandbox runner is configured — both + # execute through it and would fail on every call without one. DEFAULT_CHAT_TOOLS: list = [ "memory", "read_webpage", @@ -349,81 +307,62 @@ class Settings(BaseSettings): COMPRESSION_MODEL_OVERRIDE: Optional[str] = None # Use different model for compression COMPRESSION_PROMPT_VERSION: str = "v1.0" # Track prompt iterations COMPRESSION_MAX_HISTORY_POINTS: int = 3 # Keep only last N compression points to prevent DB bloat - COMPRESSION_RECENT_FIELD_MAX_TOKENS: int = 8000 # Per-field cap on the verbatim tail kept after a compression point (0 disables) - TOOL_RESULT_MAX_TOKENS: int = 20000 # Cap on a single tool result entering the LLM context (0 disables); journal/DB keep the full result + # Per-field cap on the verbatim tail kept after a compression point (0 disables). + COMPRESSION_RECENT_FIELD_MAX_TOKENS: int = 8000 + # Cap on one tool result entering the LLM context (0 disables); journal/DB keep it whole. + TOOL_RESULT_MAX_TOKENS: int = 20000 # Agent Guardrails - # Master switch. When False, no guardrail stage runs regardless of what an - # agent's config says. - GUARDRAILS_ENABLED: bool = True - # Registry-key allowlist; values must match GuardrailCreator.checks keys. - # Empty means "every registered check". + GUARDRAILS_ENABLED: bool = True # master switch; False disables every stage + # Allowlist of GuardrailCreator.checks keys; empty means every registered check. GUARDRAILS_CHECKS_ENABLED: list = [] - # Instance floor: a GuardrailsConfig fragment every agent inherits and - # cannot weaken. Agents may add controls or make an action stricter, never - # looser. "enabled" is required — without it the floor parses but applies - # to nothing. Example: + # A GuardrailsConfig fragment every agent inherits and cannot weaken; agents may add + # controls or make an action stricter, never looser. "enabled" is required — without it + # the floor parses but applies to nothing. Example: # {"enabled": true, "mode": "scan_all", - # "controls": [{"check": "secrets", "stage": "output", - # "action": "redact"}]} + # "controls": [{"check": "secrets", "stage": "output", "action": "redact"}]} GUARDRAILS_FLOOR: dict = {} - # Judge model for the topic/policy checks. None reuses the request's model. + # Judge model for the topic/policy checks; None reuses the request's model. GUARDRAILS_JUDGE_MODEL: Optional[str] = None - # Persist scanned text alongside guardrail_events. Off by default: the - # pre-redaction text is exactly the sensitive material a PII control exists - # to keep out of storage. + # Persist scanned text alongside guardrail_events. Off by default: pre-redaction text is + # exactly the material a PII control exists to keep out of storage. GUARDRAILS_STORE_SCANNED_TEXT: bool = False GUARDRAILS_EVENTS_RETENTION_DAYS: int = Field(default=30, ge=1) - # Internal SSE push channel (notifications + durable replay journal) - # Master switch — when False, /api/events emits a "push_disabled" comment - # and returns; clients fall back to polling. Publisher becomes a no-op. + # Internal SSE push channel (notifications + durable replay journal). + # False makes /api/events emit "push_disabled" and return; clients fall back to polling. ENABLE_SSE_PUSH: bool = True - # Per-user durable backlog cap (~entries). At typical event rates this - # gives ~24h of replay; tune up for verbose feeds, down for memory. + # Per-user durable backlog cap in entries; ~24h of replay at typical rates. EVENTS_STREAM_MAXLEN: int = 1000 - # Bounds uvicorn's graceful-shutdown drain (uvicorn_worker doesn't forward - # --graceful-timeout). Keep below the gunicorn --timeout (180) watchdog. - # Used by gunicorn_worker.BoundedDrainUvicornWorker. + # Bounds uvicorn's shutdown drain (uvicorn_worker doesn't forward --graceful-timeout). + # Keep below the gunicorn --timeout (180) watchdog. Used by BoundedDrainUvicornWorker. GRACEFUL_SHUTDOWN_TIMEOUT_SECONDS: int = 30 WSGI_THREADPOOL_WORKERS: int = 96 SSE_KEEPALIVE_SECONDS: int = Field(default=15, ge=1) - # Cap on simultaneous SSE connections per user. Each connection holds - # one WSGI thread (32 per gunicorn worker) and one Redis pub/sub - # connection. 8 covers normal multi-tab use without letting one user - # starve the pool. Set to 0 to disable the cap. + # Simultaneous SSE connections per user; each holds a WSGI thread and a Redis pub/sub + # connection. 8 covers multi-tab use without one user starving the pool. 0 disables. SSE_MAX_CONCURRENT_PER_USER: int = 8 - # Per-request cap on the number of backlog entries XRANGE returns - # for ``/api/events`` snapshots. Bounds the bytes a single replay - # can move from Redis to the wire — a malicious client looping - # ``Last-Event-ID=`` reconnects can only enumerate this - # many entries per round-trip. Combined with the per-user - # connection cap above and the windowed budget below, total - # enumeration throughput is bounded. + # Backlog entries XRANGE returns per /api/events snapshot. Bounds what one replay moves + # from Redis to the wire: a client looping Last-Event-ID reconnects enumerates at most + # this many per round-trip, and the budget below bounds total throughput. EVENTS_REPLAY_MAX_PER_REQUEST: int = 200 EVENTS_REPLAY_MAX_AGE_HOURS: int = 48 - # Sliding-window cap on snapshot replays per user. Once the budget - # is exhausted the route returns HTTP 429 with the cursor pinned; - # the client backs off and retries after the window rolls over. + # Sliding-window cap on snapshot replays per user; exhausting it returns 429 with the + # cursor pinned so the client backs off until the window rolls over. EVENTS_REPLAY_BUDGET_REQUESTS_PER_WINDOW: int = 30 EVENTS_REPLAY_BUDGET_WINDOW_SECONDS: int = 60 - # Retention for the ``message_events`` journal. The ``cleanup_message_events`` - # beat task deletes rows older than this. Reconnect-replay only - # needs the journal for streams a client could still be tailing, - # so 14 days is a generous default that covers paused/tool-action - # flows without unbounded table growth. + # Retention for the message_events journal, enforced by the cleanup_message_events beat + # task. Replay only needs streams a client could still be tailing. MESSAGE_EVENTS_RETENTION_DAYS: int = 14 # Remote Device feature. REMOTE_DEVICE_SESSION_IDLE_SECONDS: int = 60 REMOTE_DEVICE_REQUIRE_SIGNATURE: bool = False REMOTE_DEVICE_PAIRING_TTL_SECONDS: int = 600 - # Redis-backed broker tunables (route invocations cross-process so a - # scheduled/Celery run reaches the web-held device session). The command - # queue TTL must exceed the max command drain deadline (the tool caps - # timeout_ms at 600s, drained with a +5s margin = 605s) so a queued command - # for a briefly-offline device isn't evicted before its own drain gives up. + # Redis broker tunables, routing invocations cross-process so a scheduled run reaches the + # web-held device session. The queue TTL must exceed the max drain deadline (605s) so a + # command for a briefly-offline device isn't evicted before its own drain gives up. REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS: int = 900 REMOTE_DEVICE_INVOCATION_TTL_SECONDS: int = 900 REMOTE_DEVICE_OUTPUT_STREAM_MAXLEN: int = 10_000 @@ -438,23 +377,19 @@ class Settings(BaseSettings): SCHEDULE_ONCE_MAX_HORIZON: int = 31_536_000 SCHEDULE_RUN_OUTPUT_RETENTION_DAYS: int = 90 - # Code-execution sandbox (see artifacts-code-execution-spec.md §4 C2). - # The app is a CLIENT of an always-on runner; defaults are safe so app - # import never fails when the sandbox is unconfigured. + # Code-execution sandbox. The app is a CLIENT of an always-on runner; defaults are safe so + # app import never fails when the sandbox is unconfigured. SANDBOX_BACKEND: str = "jupyter" # "jupyter" (self-host) | "daytona" (Daytona Cloud) # URL of the Jupyter Kernel Gateway runner (the docsgpt-sandbox service). SANDBOX_GATEWAY_URL: str = "http://localhost:8888" SANDBOX_GATEWAY_AUTH_TOKEN: Optional[str] = None # gateway auth token, if set - # Kernelspec launched per session. Defaults to the env-scrubbing "docsgpt-python" - # spec (shipped by the docsgpt-sandbox runner) so kernel code cannot read the - # gateway auth token or operator secrets from os.environ. The stock "python3" - # spec inherits the gateway env verbatim and must not be used with untrusted code. + # Kernelspec per session. The env-scrubbing "docsgpt-python" spec keeps kernel code from + # reading the gateway token or operator secrets from os.environ; the stock "python3" spec + # inherits the gateway env verbatim and must not be used with untrusted code. SANDBOX_KERNEL_NAME: str = "docsgpt-python" SANDBOX_MAX_TTL: int = 1200 # hard cap (s) on agent-selectable keep-alive TTL - # Per-process/worker cap on concurrent live sandbox sessions. Backend-agnostic - # (complements DAYTONA_MAX_SANDBOXES); when reached, an LRU-idle session is - # evicted to make room. This bound is local to each app/worker process. - # 0 (or any non-positive value) disables the cap (unlimited sessions). + # Concurrent live sessions per process, backend-agnostic; at the cap an LRU-idle session is + # evicted. 0 or negative disables the cap. SANDBOX_MAX_SESSIONS: int = 32 SANDBOX_EXEC_TIMEOUT: int = 60 # default wall-clock cap (s) per exec call SANDBOX_HTTP_TIMEOUT: int = 10 # fixed cap (s) for REST control calls (create/delete/alive/interrupt) @@ -464,44 +399,31 @@ class Settings(BaseSettings): # ``read_document`` parsing on a dedicated Celery ``parsing`` queue (backend parser). DOCUMENT_PARSE_QUEUE: str = "parsing" # queue the parse_document task is routed to DOCUMENT_PARSE_TIMEOUT: int = 120 # seconds the tool awaits the enqueued parse before degrading - # The base timeout is a FLOOR: the awaited window (and the task's per-call Celery - # time limits) grow with the document's size, because OCR cost scales with pages - # -- a 30-page scan needs ~60s, a 100-page scan several minutes. Without this a - # large scan is silently dropped from the node/tool at the base window. + # The base timeout is a FLOOR: the window grows with document size, because OCR cost scales + # with pages. Without this a large scan is silently dropped at the base window. DOCUMENT_PARSE_TIMEOUT_PER_MB: int = 60 # extra seconds of parse window per MiB of input DOCUMENT_PARSE_TIMEOUT_MAX: int = 900 # absolute ceiling on the size-scaled parse window DOCUMENT_PARSE_MAX_BYTES: int = 0 # cap on a parsed document's bytes (0 = reuse SANDBOX_MAX_INPUT_BYTES) DOCUMENT_MAX_DECOMPRESSED_BYTES: int = 300 * 1024 * 1024 DOCUMENT_MAX_ARCHIVE_ENTRIES: int = 10000 - # Per-agent-node cap on files passed natively to the node's LLM (vision/doc - # inputs). Files past the cap are extracted to text or dropped, not attached - # natively, to bound context/cost. Re-uses SANDBOX_MAX_INPUT_BYTES per file. + # Files per node passed natively to the LLM; past the cap they are extracted to text or + # dropped, to bound context and cost. Re-uses SANDBOX_MAX_INPUT_BYTES per file. WORKFLOW_NODE_NATIVE_MAX_FILES: int = 5 - # Per-agent-node cap on documents extracted to text via the parsing worker. - # Each non-native, non-text document issues a separate blocking parse, so a - # node referencing many documents (e.g. the ``*`` token) is bounded here to - # avoid serializing dozens of parses; documents past the cap are skipped with - # a truncation note instead of extracted. + # Documents per node extracted via the parsing worker. Each issues a separate blocking + # parse; past the cap they are skipped with a truncation note. WORKFLOW_NODE_EXTRACT_MAX_FILES: int = 5 - # Total wall clock one node may spend on blocking document parses, shared - # across all of them. The per-document window scales with size (up to - # DOCUMENT_PARSE_TIMEOUT_MAX), so without a shared budget a node could - # serialize WORKFLOW_NODE_EXTRACT_MAX_FILES full windows and hold a web - # threadpool slot for that whole time. + # Wall clock one node may spend on blocking parses, shared across all of them. Without it a + # node could serialize WORKFLOW_NODE_EXTRACT_MAX_FILES full windows on a web threadpool slot. WORKFLOW_NODE_EXTRACT_BUDGET_SECONDS: int = 900 - # A workflow run row is pre-created as ``running`` and finalized when its - # generator completes; a client disconnect or worker crash can strand it in - # ``running`` forever. The beat reaper fails runs still ``running`` past this - # many seconds. Generous so a legitimately long run is never cut off. + # A run row is pre-created as ``running``; a disconnect or crash can strand it there. The + # beat reaper fails runs still ``running`` past this. Generous so a long run is never cut off. WORKFLOW_RUN_STALE_SECONDS: int = 3600 - # Runner container resource caps — consumed by the docsgpt-sandbox compose - # service (deployment/sandbox), not by the app client. cgroup CPU/mem caps - # are part of the untrusted-code security boundary. + # Runner container caps, consumed by the docsgpt-sandbox compose service, not the app. + # These cgroup limits are part of the untrusted-code security boundary. SANDBOX_MEMORY: str = "1g" # docker mem_limit for the runner container SANDBOX_CPUS: str = "1.0" # docker cpu quota for the runner container - # Daytona Cloud managed backend (used only when SANDBOX_BACKEND="daytona"). - # The app is a REST client of Daytona Cloud authenticated by DAYTONA_API_KEY; - # all knobs are optional so app import never fails when the backend is unused. + # Daytona Cloud backend (SANDBOX_BACKEND="daytona"). All knobs are optional so app import + # never fails when the backend is unused. DAYTONA_API_KEY: Optional[str] = None # Daytona Cloud API key (secret) DAYTONA_API_URL: Optional[str] = None # override Daytona API base URL, if self-targeting DAYTONA_TARGET: Optional[str] = None # Daytona region/target, e.g. "us" @@ -510,8 +432,7 @@ class Settings(BaseSettings): DAYTONA_AUTO_STOP_INTERVAL: int = 15 # minutes idle before Daytona auto-stops a sandbox (0 disables) DAYTONA_AUTO_DELETE_INTERVAL: int = 60 # minutes after stop before Daytona auto-deletes (-1 disables) DAYTONA_MAX_SANDBOXES: int = 50 # cap on concurrent live Daytona sandboxes (cost-DoS guard) - # Per-user artifact quotas (generous defaults; enforced at persistence time). - # For all three, 0 (or any non-positive value) disables that quota (unlimited). + # Per-user artifact quotas, enforced at persistence time. 0 or negative disables a quota. ARTIFACT_MAX_BYTES: int = 50 * 1024 * 1024 # cap on a single stored artifact version's bytes ARTIFACT_MAX_COUNT_PER_USER: int = 5000 # cap on artifacts a user may own ARTIFACT_MAX_TOTAL_BYTES_PER_USER: int = 5 * 1024 * 1024 * 1024 # cap on a user's total stored bytes diff --git a/application/scripts/reembed.py b/application/scripts/reembed.py index 0434d21d..bd325771 100644 --- a/application/scripts/reembed.py +++ b/application/scripts/reembed.py @@ -411,6 +411,17 @@ def main(argv: Optional[Sequence[str]] = None) -> int: args = build_parser().parse_args(argv) _log_setup(args.verbose) + # 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 + # embeds every chunk in the index, where a broker round trip per batch adds + # latency and a dependency on a worker running. Loading the model here also + # means the script reports a real failure for a model it cannot load, + # instead of timing out against an empty queue. + if getattr(settings, "EMBEDDINGS_DELEGATE_TO_WORKER", False): + logger.info("Embedding in-process; worker delegation does not apply here.") + settings.EMBEDDINGS_DELEGATE_TO_WORKER = False + store_type = (settings.VECTOR_STORE or "").lower() if store_type not in SUPPORTED_STORES: logger.error( diff --git a/application/storage/db/bootstrap.py b/application/storage/db/bootstrap.py index 19dce2d1..758f5227 100644 --- a/application/storage/db/bootstrap.py +++ b/application/storage/db/bootstrap.py @@ -127,21 +127,30 @@ def ensure_vector_schema(*, logger: Optional[logging.Logger] = None) -> None: PGVectorStore, ) - # Loading the model here is deliberate: the process loads it on first - # retrieval anyway, and EmbeddingsSingleton caches it. - dim: Optional[int] = None - try: - from application.vectorstore.base import get_embeddings + # All this needs is an integer, and for a model the registry describes that + # is a lookup. It used to construct the embeddings instance, which loaded + # ~800 MB of ONNX into every API and worker process at import purely to + # read ``.dimension`` off it. + from application.vectorstore.model_registry import dimension_for - dim = getattr(get_embeddings(), "dimension", None) - except Exception as exc: # noqa: BLE001 — never block boot on the model - log.warning( - "ensure_vector_schema: could not load the embeddings model (%s); " - "creating the table with %d dimensions and skipping the dimension " - "check.", - exc, - DEFAULT_EMBEDDING_DIM, - ) + 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 + + dim = getattr(build_local_embeddings(), "dimension", None) + except Exception as exc: # noqa: BLE001 — never block boot on the model + log.warning( + "ensure_vector_schema: could not load the embeddings model (%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; " diff --git a/application/vectorstore/base.py b/application/vectorstore/base.py index f503d9fd..b1899380 100644 --- a/application/vectorstore/base.py +++ b/application/vectorstore/base.py @@ -270,6 +270,16 @@ def _azure_configured() -> bool: ) +def _delegation_enabled() -> bool: + """True when this process should embed on the worker rather than locally. + + Compared against ``True`` rather than coerced: tests patch ``settings`` + with a ``MagicMock``, whose every attribute is a truthy object, and + ``bool()`` on that would silently route them through the broker. + """ + return getattr(settings, "EMBEDDINGS_DELEGATE_TO_WORKER", False) is True + + def get_embeddings( embeddings_name: Optional[str] = None, embeddings_key: Optional[str] = None ): @@ -279,6 +289,11 @@ def get_embeddings( skip the remote dispatch and the OpenAI/Azure key handling. Route every caller through here. + With ``EMBEDDINGS_DELEGATE_TO_WORKER`` this returns a client that runs the + model on the Celery worker, so an API process never loads one. The client + embeds locally when it finds itself inside a worker task, so the worker is + unaffected. + Args: embeddings_name: Model name; defaults to ``settings.EMBEDDINGS_NAME``. embeddings_key: API key; defaults to ``settings.EMBEDDINGS_KEY``. @@ -287,6 +302,34 @@ def get_embeddings( The shared embeddings instance for the resolved model. """ embeddings_name = embeddings_name or settings.EMBEDDINGS_NAME + if not settings.EMBEDDINGS_BASE_URL and _delegation_enabled(): + cache_key = f"delegated_{embeddings_name}" + if cache_key not in EmbeddingsSingleton._instances: + from application.vectorstore.embeddings_delegated import DelegatedEmbeddings + + EmbeddingsSingleton._instances[cache_key] = DelegatedEmbeddings( + embeddings_name, embeddings_key + ) + return EmbeddingsSingleton._instances[cache_key] + return build_local_embeddings(embeddings_name, embeddings_key) + + +def build_local_embeddings( + embeddings_name: Optional[str] = None, embeddings_key: Optional[str] = None +): + """Resolve the embeddings instance that runs in *this* process. + + Bypasses worker delegation, so it is what the worker's embed task and the + boot hook use. Everything else should call :func:`get_embeddings`. + + Args: + embeddings_name: Model name; defaults to ``settings.EMBEDDINGS_NAME``. + embeddings_key: API key; defaults to ``settings.EMBEDDINGS_KEY``. + + Returns: + The shared in-process embeddings instance for the resolved model. + """ + embeddings_name = embeddings_name or settings.EMBEDDINGS_NAME embeddings_key = ( embeddings_key if embeddings_key is not None else settings.EMBEDDINGS_KEY ) diff --git a/application/vectorstore/embeddings_delegated.py b/application/vectorstore/embeddings_delegated.py new file mode 100644 index 00000000..b952537a --- /dev/null +++ b/application/vectorstore/embeddings_delegated.py @@ -0,0 +1,118 @@ +"""Query embedding executed in the Celery worker instead of in the API. + +The API embeds every query it serves, so it needs an embedder -- and a local +one costs roughly 800 MB of ONNX Runtime per process. That is the whole +footprint of an API container that otherwise holds no model. + +This client keeps the interface (``embed_query``/``embed_documents``/ +``dimension``) and moves only the computation: the text goes to the worker over +Celery and the vector comes back. The API pays a broker round trip per query +and no resident model. + +Inside a worker there is nothing to delegate to -- dispatching would queue work +behind the task already running and wait on itself -- so a call made while a +task is executing runs locally, on a model this process loads once and caches. +``DOCUMENT_PARSE_QUEUE`` exists for the same reason on the parsing side. + +Production deployments should point ``EMBEDDINGS_BASE_URL`` at a real embedding +service instead: that removes the model from *both* processes and costs a +network hop rather than a broker round trip. +""" + +from __future__ import annotations + +import logging +from typing import Any, List, Optional + +from application.core.settings import settings +from application.vectorstore.model_registry import dimension_for + +logger = logging.getLogger(__name__) + +#: Dispatched by name so the API never imports the task module -- and through +#: it ``application.worker``, which pulls in the whole parsing stack. +EMBED_TASK = "application.vectorstore.embeddings_tasks.embed_texts" + + +def _in_worker() -> bool: + """True when a Celery task is executing in this process.""" + try: + from application.celery_init import celery + + return celery.current_worker_task is not None + except Exception: + return False + + +class DelegatedEmbeddings: + """Embeds by dispatching to the Celery worker, or locally inside one.""" + + def __init__(self, embeddings_name: str, embeddings_key: Optional[str] = None) -> None: + self.embeddings_name = embeddings_name + self.embeddings_key = embeddings_key + self._local: Any = None + self._dimension: Optional[int] = dimension_for(embeddings_name) + + def _local_embeddings(self): + """The in-process model, built once, for use inside a worker task.""" + if self._local is None: + from application.vectorstore.base import build_local_embeddings + + self._local = build_local_embeddings(self.embeddings_name, self.embeddings_key) + return self._local + + def _dispatch(self, texts: List[str]) -> List[List[float]]: + """Run the embed task on the worker and wait for its vectors.""" + from application.celery_init import celery + + queue = getattr(settings, "EMBEDDINGS_QUEUE", "embeddings") + timeout = getattr(settings, "EMBEDDINGS_DELEGATE_TIMEOUT", 60) + result = celery.send_task(EMBED_TASK, args=[texts, self.embeddings_name], queue=queue) + try: + return result.get(timeout=timeout) + except Exception as exc: + 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." + ) from exc + + def embed_documents(self, documents: List[str]) -> List[List[float]]: + """Embed a list of texts, preserving order.""" + if not documents: + return [] + if _in_worker(): + return self._local_embeddings().embed_documents(documents) + vectors = self._dispatch(list(documents)) + if self._dimension is None and vectors: + self._dimension = len(vectors[0]) + return vectors + + def embed_query(self, query: str) -> List[float]: + """Embed a single query string.""" + return self.embed_documents([query])[0] + + @property + def dimension(self) -> Optional[int]: + """Vector width, from the registry where possible. + + Falls back to one round trip for a model the registry does not + describe, and to ``None`` when even that fails -- callers already treat + an unknown width as "nothing to compare yet" rather than an error. + """ + if self._dimension is None: + try: + self._dimension = len(self.embed_query("dimension probe")) + except Exception as exc: + logger.warning("Could not determine embedding width: %s", exc) + return None + return self._dimension + + def __call__(self, text): + if isinstance(text, str): + return self.embed_query(text) + elif isinstance(text, list): + return self.embed_documents(text) + raise ValueError("Input must be a string or a list of strings") diff --git a/application/vectorstore/embeddings_local.py b/application/vectorstore/embeddings_local.py index 10cb68f0..f35854c5 100644 --- a/application/vectorstore/embeddings_local.py +++ b/application/vectorstore/embeddings_local.py @@ -312,20 +312,21 @@ class EmbeddingsWrapper: def embed_documents(self, documents: List[str]) -> List[List[float]]: """Embed a list of documents, preserving input order. - Inputs are grouped by length before batching. ONNX needs a rectangular - tensor, so every input in a batch is padded up to the longest one in - it; with mixed lengths that padding is most of the work. Grouping - similar lengths together measured 19% faster and 45% lower peak memory - on production-sized chunks, and at full context an unsorted batch of 4 - was slower than no batching at all. + Batched by ``EMBEDDINGS_MODEL_BATCH_SIZE``, not by the pipeline's + ``EMBEDDINGS_BATCH_SIZE``: one is documents per forward pass, the other + is chunks per store transaction, and sizing the forward pass from the + transaction is what made ingest peak at 6.6 GB. - The original order is restored before returning, so callers zipping - these against their texts are unaffected. + Inputs are grouped by length first. ONNX needs a rectangular tensor, so + every input in a pass is padded up to the longest one in it; with mixed + lengths that padding is most of the work. The original order is + restored before returning, so callers zipping these against their texts + are unaffected. """ if not documents: return [] batch_size: Optional[int] = None - raw = getattr(settings, "EMBEDDINGS_BATCH_SIZE", None) + raw = getattr(settings, "EMBEDDINGS_MODEL_BATCH_SIZE", None) if isinstance(raw, int) and not isinstance(raw, bool) and raw > 0: batch_size = raw diff --git a/application/vectorstore/embeddings_tasks.py b/application/vectorstore/embeddings_tasks.py new file mode 100644 index 00000000..82f142c3 --- /dev/null +++ b/application/vectorstore/embeddings_tasks.py @@ -0,0 +1,29 @@ +"""The Celery task behind :mod:`application.vectorstore.embeddings_delegated`. + +Kept out of ``application.api.user.tasks`` deliberately: that module imports +``application.worker`` and the whole parsing stack with it, which is the +opposite of what delegation is for. +""" + +from __future__ import annotations + +from typing import List, Optional + +from application.celery_init import celery +from application.vectorstore.embeddings_delegated import EMBED_TASK + + +@celery.task(name=EMBED_TASK, acks_late=False, ignore_result=False) +def embed_texts(texts: List[str], embeddings_name: Optional[str] = None) -> List[List[float]]: + """Embed ``texts`` with the worker's local model. + + Args: + texts: Strings to embed. + embeddings_name: Model to use; the configured one when omitted. + + Returns: + One vector per input, in input order. + """ + from application.vectorstore.base import get_embeddings + + return get_embeddings(embeddings_name).embed_documents(list(texts)) diff --git a/docs/content/Deploying/Development-Environment.mdx b/docs/content/Deploying/Development-Environment.mdx index b94e8056..7f78d719 100644 --- a/docs/content/Deploying/Development-Environment.mdx +++ b/docs/content/Deploying/Development-Environment.mdx @@ -76,18 +76,16 @@ To run the DocsGPT backend locally, you'll need to set up a Python environment a venv/Scripts/activate ``` -3. **Download Embedding Model:** +3. **Embedding Model (no action needed):** - The backend requires an embedding model. Download the `mpnet-base-v2` model and place it in the `models/` directory within the project root. You can use the following script: + The embedding model is downloaded automatically the first time you ingest a document, and cached for subsequent runs. Set `EMBEDDINGS_CACHE_DIR` to control where. + + For an offline or air-gapped machine, fetch it ahead of time instead: ```bash - wget https://d3dg1063dc54p9.cloudfront.net/models/embeddings/mpnet-base-v2.zip - unzip mpnet-base-v2.zip -d model - rm mpnet-base-v2.zip + python -m application.scripts.prefetch_models ``` - Alternatively, you can manually download the zip file from [here](https://d3dg1063dc54p9.cloudfront.net/models/embeddings/mpnet-base-v2.zip), unzip it, and place the extracted folder in `models/`. - 4. **Install Backend Dependencies:** Navigate to the root of your DocsGPT repository and install the required Python packages: diff --git a/docs/content/Deploying/DocsGPT-Settings.mdx b/docs/content/Deploying/DocsGPT-Settings.mdx index b71c93d6..84b3bc6e 100644 --- a/docs/content/Deploying/DocsGPT-Settings.mdx +++ b/docs/content/Deploying/DocsGPT-Settings.mdx @@ -59,8 +59,9 @@ Here are some of the most fundamental settings you'll likely want to configure: - **`EMBEDDINGS_NAME`**: This setting defines which embedding model DocsGPT will use to generate vector embeddings for your documents. Embeddings are numerical representations of text that allow DocsGPT to understand the semantic meaning of your documents for efficient search and retrieval. - - **Default value:** `huggingface_sentence-transformers/all-mpnet-base-v2` (a good general-purpose embedding model). - - **Other options:** You can explore other embedding models from Hugging Face Sentence Transformers or other providers if needed. + - **Default value:** `huggingface_sentence-transformers/all-mpnet-base-v2`, kept so an existing index stays readable. New installs are pointed at `ibm-granite/granite-embedding-311m-multilingual-r2` (multilingual, 32k context, same 768 dimensions) by `.env-template` and `setup.sh`. + - **Other options:** Any FastEmbed built-in model, or any Hugging Face repository shipping an ONNX export. See [Embeddings](/Models/embeddings). + - **Changing it on an existing index requires re-embedding** — same-width models swap without any error and silently degrade retrieval. Run `python -m application.scripts.reembed`. - **`API_KEY`**: Required for most cloud-based LLM providers. This is your authentication key to access the LLM provider's API. You'll need to obtain this key from your chosen provider's platform. @@ -93,7 +94,7 @@ LLM_PROVIDER=openai # Using OpenAI compatible API format for local models API_KEY=None # API Key is not needed for local Ollama LLM_NAME=llama3.2:1b OPENAI_BASE_URL=http://host.docker.internal:11434/v1 # Default Ollama API URL within Docker -EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2 # You can also run embeddings locally if needed +EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2 # runs locally; see Models/embeddings for alternatives ``` In this case, even though you are using Ollama locally, `LLM_PROVIDER` is set to `openai` because Ollama (and many other local inference engines) are designed to be API-compatible with OpenAI. `OPENAI_BASE_URL` points DocsGPT to the local Ollama server. @@ -442,7 +443,7 @@ See [Embeddings](/Models/embeddings) for full guidance. | Setting | Default | Description | | --- | --- | --- | -| `EMBEDDINGS_NAME` | `huggingface_sentence-transformers/all-mpnet-base-v2` | The embedding model. | +| `EMBEDDINGS_NAME` | `huggingface_sentence-transformers/all-mpnet-base-v2` | The embedding model. New installs use `ibm-granite/granite-embedding-311m-multilingual-r2`. Changing it on a populated index requires `application.scripts.reembed`. | | `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). | diff --git a/docs/content/Models/embeddings.md b/docs/content/Models/embeddings.md index 634eaaee..ba15a5e1 100644 --- a/docs/content/Models/embeddings.md +++ b/docs/content/Models/embeddings.md @@ -24,36 +24,30 @@ In essence, embedding models are the bridge that allows DocsGPT to understand th DocsGPT is designed to be flexible and supports a wide range of embedding models right out of the box: -* **Sentence Transformers:** DocsGPT supports all models available through the [Sentence Transformers library](https://www.sbert.net/). This library offers a vast selection of pre-trained embedding models, known for their quality and efficiency in various semantic tasks. This is the default (`EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2`). +* **Local models (FastEmbed / ONNX Runtime):** DocsGPT runs local embeddings through [FastEmbed](https://github.com/qdrant/fastembed). Any of FastEmbed's built-in models works, as does any Hugging Face repository that ships an ONNX export at `onnx/model.onnx` — which covers most popular sentence-transformers repos. A repository with PyTorch weights only will not load; serve those over `EMBEDDINGS_BASE_URL` instead. New installs default to `ibm-granite/granite-embedding-311m-multilingual-r2`; existing ones stay on `huggingface_sentence-transformers/all-mpnet-base-v2` until re-embedded. * **OpenAI Embeddings:** DocsGPT supports OpenAI embedding models (for example `text-embedding-ada-002`, `text-embedding-3-small`, `text-embedding-3-large`) via the OpenAI API. * **Azure OpenAI Embeddings:** Set `AZURE_EMBEDDINGS_DEPLOYMENT_NAME` alongside your Azure OpenAI configuration. * **Remote OpenAI-compatible Embeddings:** Any server that exposes an OpenAI-compatible `/v1/embeddings` endpoint (for example llama.cpp, vLLM, TEI, or a hosted provider) by setting `EMBEDDINGS_BASE_URL`. See [Remote Embeddings](#remote-openai-compatible-embeddings) below. -## Configuring Sentence Transformer Models +## Configuring a Local Model -To utilize Sentence Transformer models within DocsGPT, you need to follow these steps: +Set `EMBEDDINGS_NAME` in your `.env` to a registry name or a Hugging Face repository id: -1. **Download the Model:** Sentence Transformer models are typically hosted on Hugging Face Model Hub. You need to download your chosen model and place it in the `model/` folder in the root directory of your DocsGPT project. +``` +EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2 +``` - For example, to use the `all-mpnet-base-v2` model, you would set `EMBEDDINGS_NAME` as described below, and ensure that the model files are available locally (DocsGPT will attempt to download it if it's not found, but local download is recommended for development and offline use). +The model is downloaded on first use and cached; set `EMBEDDINGS_CACHE_DIR` to control where. There is no `model/` folder to populate by hand, and a filesystem path is not accepted as a model name. -2. **Set `EMBEDDINGS_NAME` in `.env` (or `settings.py`):** You need to configure the `EMBEDDINGS_NAME` setting in your `.env` file (or `settings.py`) to point to the desired Sentence Transformer model. +DocsGPT knows the pooling, vector width and context window of the models in its registry (`all-mpnet-base-v2`, `granite-embedding-311m-multilingual-r2`, `granite-embedding-97m-multilingual-r2`). For any other repository it reads those from the repository's own `1_Pooling/config.json` and `modules.json`. If a repository declares neither, mean pooling with L2 normalization is assumed and a warning is logged — pin the real values with `EMBEDDINGS_POOLING` (`cls` or `mean`) and `EMBEDDINGS_NORMALIZE`. - * **Using a pre-downloaded model from `model/` folder:** You can specify a path to the downloaded model within the `model/` directory. For instance, if you downloaded `all-mpnet-base-v2` and it's in `model/all-mpnet-base-v2`, you could potentially use a relative path like (though direct path to the model name is usually sufficient): +Models with a Dense projection layer (for example `sentence-transformers/LaBSE`) are refused at startup: FastEmbed cannot apply the projection, so the vectors would be the wrong width and in a different space. - ``` - EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2 - ``` - or simply use the model identifier: - ``` - EMBEDDINGS_NAME=sentence-transformers/all-mpnet-base-v2 - ``` +For an offline or air-gapped install, pre-fetch the model at build or setup time: - * **Using a model directly from Hugging Face Model Hub:** You can directly specify the model identifier from Hugging Face Model Hub: - - ``` - EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2 - ``` +```bash +python -m application.scripts.prefetch_models +``` ## Using OpenAI Embeddings @@ -95,6 +89,41 @@ You usually do not need to set it. When `EMBEDDINGS_NAME` names a model DocsGPT Leaving `EMBEDDINGS_NAME` unset imposes no limit: the name is only forwarded as the `model` field in each request, so a default nobody chose is not taken as a description of your server. When no model tokenizer is available, counting falls back to tiktoken — pick a limit with headroom below the server's true limit to absorb the skew between the two tokenizers. +## Where the model runs + +A local embedding model costs a few hundred megabytes of resident memory per process, and the API embeds every query it serves — so by default it would hold its own copy alongside the worker's. + +`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: + +```bash +celery -A application.app.celery worker -Q embeddings +``` + +Set `EMBEDDINGS_DELEGATE_TO_WORKER=false` if you run the API without a worker; it will load the model in-process instead. + +For production, prefer `EMBEDDINGS_BASE_URL`. A real embedding service removes the model from *both* the API and the worker, and replaces the broker round trip with a network call. + +## Batch sizes + +Two separate knobs, easily confused: + +- `EMBEDDINGS_BATCH_SIZE` (default 32) — chunks per store transaction, and per request to a remote embeddings API. Larger means fewer round trips and fewer transactions. +- `EMBEDDINGS_MODEL_BATCH_SIZE` (default 1) — documents per forward pass of a *local* model. + +For the local model, bigger batches are not faster. ONNX needs a rectangular tensor, so every input in a pass is padded up to the longest one in it, and that waste grows with the square of chunk length. Measured on a 30-document ingest at the 1250-token default chunk size: + +| `EMBEDDINGS_MODEL_BATCH_SIZE` | embed time | peak RSS | +| --- | --- | --- | +| 32 | 154 s | 7.7 GB | +| 8 | 76 s | 5.0 GB | +| 4 | 76 s | 3.6 GB | +| 2 | 74 s | 2.3 GB | +| 1 | 53 s | 1.5 GB | + +Raise it only if your chunks are short and uniform in length. + ## Important: Embedding Dimensions Must Stay Consistent Each embedding model produces vectors of a fixed dimension, and your vector store is created with that dimension. **Changing `EMBEDDINGS_NAME` to a model with a different dimension is not compatible with an existing index** — FAISS and LanceDB will raise a dimension-mismatch error, and pgvector/Qdrant tables are sized to the original dimension. @@ -119,6 +148,6 @@ With `GRAPHRAG_ENABLED`, the script also rewrites `graph_nodes.name_embedding` o ## Adding Support for Other Embedding Models -If you wish to use an embedding model that is not supported out-of-the-box, a good starting point for adding custom embedding model support is to examine the `base.py` file located in the `application/vectorstore` directory. +To teach DocsGPT about a new model — so it carries a known pooling, width and context window rather than being inferred — add an `EmbeddingModel` entry to `MODELS` in `application/vectorstore/model_registry.py`. That registry is the single source of truth the local runner, the remote client, the schema bootstrap and the chunker all read. Specifically, pay attention to the `EmbeddingsWrapper` and `EmbeddingsSingleton` classes. `EmbeddingsWrapper` provides a way to wrap different embedding model libraries into a consistent interface for DocsGPT. `EmbeddingsSingleton` manages the instantiation and retrieval of embedding model instances. By understanding these classes and the existing embedding model implementations, you can create your own custom integration for virtually any embedding model library you desire. \ No newline at end of file diff --git a/tests/conftest.py b/tests/conftest.py index 5682e183..b3b922be 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -158,6 +158,21 @@ def _no_real_redis(monkeypatch): """ monkeypatch.setattr("application.cache._redis_instance", None) monkeypatch.setattr("application.cache._redis_creation_failed", True) + + +@pytest.fixture(autouse=True) +def _no_worker_delegation(monkeypatch): + """Embed in-process during tests, the way CI has no worker to embed on. + + ``EMBEDDINGS_DELEGATE_TO_WORKER`` ships on, so an unmocked embed would + publish to a broker nobody is consuming and block for + ``EMBEDDINGS_DELEGATE_TIMEOUT`` before failing -- a minute per call, and a + pass/fail that depends on whether the developer happens to have a worker + running. Tests covering delegation patch the setting back on themselves. + """ + from application.core.settings import settings + + monkeypatch.setattr(settings, "EMBEDDINGS_DELEGATE_TO_WORKER", False, raising=False) monkeypatch.setattr("application.cache._pubsub_redis_instance", None) monkeypatch.setattr("application.cache._pubsub_redis_creation_failed", True) diff --git a/tests/scripts/test_reembed.py b/tests/scripts/test_reembed.py index 14d8e374..0790ab14 100644 --- a/tests/scripts/test_reembed.py +++ b/tests/scripts/test_reembed.py @@ -374,3 +374,19 @@ class TestGraphNodeReembedding: for call in cursor.executemany.call_args_list if "graph_nodes" in str(call.args[0]) ] + + +class TestEmbedsInProcess: + """A batch job should not round-trip every chunk through the broker.""" + + def test_delegation_is_turned_off_for_the_run(self, monkeypatch): + from application.core.settings import settings + + monkeypatch.setattr(settings, "EMBEDDINGS_DELEGATE_TO_WORKER", True, raising=False) + monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector", raising=False) + seen = {} + with patch.object(reembed, "run", side_effect=lambda *a, **k: seen.setdefault( + "delegating", settings.EMBEDDINGS_DELEGATE_TO_WORKER + ) or 0): + assert reembed.main(["--dry-run"]) == 0 + assert seen["delegating"] is False diff --git a/tests/storage/db/test_bootstrap_vector_schema.py b/tests/storage/db/test_bootstrap_vector_schema.py index ce925786..64d50f3c 100644 --- a/tests/storage/db/test_bootstrap_vector_schema.py +++ b/tests/storage/db/test_bootstrap_vector_schema.py @@ -74,8 +74,8 @@ class TestEnsureVectorSchemaCreates: cursor = MagicMock() conn.cursor.return_value = cursor with patch("psycopg.connect", return_value=conn) as connect, patch( - "application.vectorstore.base.get_embeddings", - return_value=_embeddings(dimension), + "application.vectorstore.model_registry.dimension_for", + return_value=dimension, ), patch( "application.vectorstore.pgvector.PGVectorStore.create_schema" ) as vector_schema, patch( @@ -130,8 +130,8 @@ class TestEnsureVectorSchemaDimensionCheck: ): conn = MagicMock() with patch("psycopg.connect", return_value=conn), patch( - "application.vectorstore.base.get_embeddings", - return_value=_embeddings(1536), + "application.vectorstore.model_registry.dimension_for", + return_value=1536, ), patch("application.vectorstore.pgvector.PGVectorStore.create_schema"), patch( "application.vectorstore.pgvector.PGVectorStore.table_dimension", return_value=768, @@ -226,3 +226,42 @@ class TestBootGating: assert 'os.environ.setdefault("AUTO_VECTOR_SCHEMA", "false")' in ( conftest.read_text() ) + + +@pytest.mark.unit +class TestBootDoesNotLoadTheModel: + """The hook needs an integer, not an inference session. + + It used to build the embeddings instance to read ``.dimension`` off it, + loading several hundred MB of ONNX into every API and worker process at + import. For a model the registry describes that is a lookup. + """ + + def _run(self, registry_dim, loader): + conn = MagicMock() + conn.cursor.return_value = MagicMock() + with patch("psycopg.connect", return_value=conn), patch( + "application.vectorstore.model_registry.dimension_for", + return_value=registry_dim, + ), patch( + "application.vectorstore.base.build_local_embeddings", loader + ), patch( + "application.vectorstore.pgvector.PGVectorStore.create_schema" + ) as vector_schema, patch( + "application.vectorstore.pgvector.PGVectorStore.table_dimension", + return_value=None, + ): + ensure_vector_schema() + return vector_schema + + def test_a_registered_model_is_never_constructed(self, vector_settings): + loader = MagicMock() + vector_schema = self._run(768, loader) + loader.assert_not_called() + assert vector_schema.call_args.kwargs["dimension"] == 768 + + def test_an_unregistered_model_still_falls_back_to_loading(self, vector_settings): + loader = MagicMock(return_value=_embeddings(1024)) + vector_schema = self._run(None, loader) + loader.assert_called_once() + assert vector_schema.call_args.kwargs["dimension"] == 1024 diff --git a/tests/vectorstore/test_embeddings_delegated.py b/tests/vectorstore/test_embeddings_delegated.py new file mode 100644 index 00000000..b38b7310 --- /dev/null +++ b/tests/vectorstore/test_embeddings_delegated.py @@ -0,0 +1,130 @@ +"""Query embedding runs on the worker so the API holds no model.""" + +from unittest.mock import MagicMock, patch + +import pytest + +from application.vectorstore import base +from application.vectorstore.embeddings_delegated import EMBED_TASK, DelegatedEmbeddings + + +@pytest.fixture(autouse=True) +def _clear_singleton(): + base.EmbeddingsSingleton._instances.clear() + yield + base.EmbeddingsSingleton._instances.clear() + + +@pytest.fixture +def not_in_worker(): + with patch("application.vectorstore.embeddings_delegated._in_worker", return_value=False): + yield + + +class TestDispatch: + def test_query_is_embedded_on_the_worker(self, not_in_worker): + celery = MagicMock() + celery.send_task.return_value.get.return_value = [[0.1, 0.2, 0.3]] + with patch("application.celery_init.celery", celery): + vector = DelegatedEmbeddings("some/model").embed_query("hello") + assert vector == [0.1, 0.2, 0.3] + assert celery.send_task.call_args.args[0] == EMBED_TASK + assert celery.send_task.call_args.kwargs["args"] == [["hello"], "some/model"] + + def test_routed_to_the_embeddings_queue(self, not_in_worker): + celery = MagicMock() + celery.send_task.return_value.get.return_value = [[0.0]] + with patch("application.celery_init.celery", celery): + with patch.object(base.settings, "EMBEDDINGS_QUEUE", "embeddings"): + DelegatedEmbeddings("some/model").embed_query("hi") + assert celery.send_task.call_args.kwargs["queue"] == "embeddings" + + def test_no_worker_gives_an_actionable_error(self, not_in_worker): + celery = MagicMock() + celery.send_task.return_value.get.side_effect = TimeoutError("no worker") + with patch("application.celery_init.celery", celery): + with pytest.raises(RuntimeError) as excinfo: + DelegatedEmbeddings("some/model").embed_query("hi") + message = str(excinfo.value) + assert "EMBEDDINGS_DELEGATE_TO_WORKER=false" in message + assert "EMBEDDINGS_BASE_URL" in message + + def test_empty_input_never_reaches_the_broker(self, not_in_worker): + celery = MagicMock() + with patch("application.celery_init.celery", celery): + assert DelegatedEmbeddings("some/model").embed_documents([]) == [] + celery.send_task.assert_not_called() + + +class TestInsideAWorker: + """Dispatching from inside a task would queue work behind itself.""" + + def test_a_running_task_embeds_locally(self): + local = MagicMock() + local.embed_documents.return_value = [[1.0, 2.0]] + celery = MagicMock() + with patch("application.vectorstore.embeddings_delegated._in_worker", return_value=True): + with patch("application.vectorstore.base.build_local_embeddings", return_value=local): + with patch("application.celery_init.celery", celery): + vector = DelegatedEmbeddings("some/model").embed_query("hi") + assert vector == [1.0, 2.0] + celery.send_task.assert_not_called() + + def test_the_local_model_is_built_once(self): + local = MagicMock() + local.embed_documents.return_value = [[1.0]] + builder = MagicMock(return_value=local) + client = DelegatedEmbeddings("some/model") + with patch("application.vectorstore.embeddings_delegated._in_worker", return_value=True): + with patch("application.vectorstore.base.build_local_embeddings", builder): + client.embed_query("a") + client.embed_query("b") + builder.assert_called_once() + + +class TestDimension: + def test_registry_width_costs_no_round_trip(self): + celery = MagicMock() + with patch("application.celery_init.celery", celery): + client = DelegatedEmbeddings("ibm-granite/granite-embedding-311m-multilingual-r2") + assert client.dimension == 768 + celery.send_task.assert_not_called() + + def test_unknown_width_is_probed_once(self, not_in_worker): + celery = MagicMock() + celery.send_task.return_value.get.return_value = [[0.0] * 1024] + with patch("application.celery_init.celery", celery): + client = DelegatedEmbeddings("some/unregistered") + assert client.dimension == 1024 + assert client.dimension == 1024 + celery.send_task.assert_called_once() + + def test_an_unreachable_worker_reports_no_width(self, not_in_worker): + celery = MagicMock() + celery.send_task.return_value.get.side_effect = TimeoutError("down") + with patch("application.celery_init.celery", celery): + assert DelegatedEmbeddings("some/unregistered").dimension is None + + +class TestGetEmbeddingsDispatch: + def test_delegates_when_enabled(self): + with patch.object(base.settings, "EMBEDDINGS_BASE_URL", None): + with patch.object(base.settings, "EMBEDDINGS_DELEGATE_TO_WORKER", True): + assert isinstance(base.get_embeddings("some/model"), DelegatedEmbeddings) + + def test_remote_url_wins_over_delegation(self): + with patch.object(base.settings, "EMBEDDINGS_BASE_URL", "http://embed.local"): + with patch.object(base.settings, "EMBEDDINGS_DELEGATE_TO_WORKER", True): + assert isinstance(base.get_embeddings("some/model"), base.RemoteEmbeddings) + + def test_disabled_loads_in_process(self): + with patch.object(base.settings, "EMBEDDINGS_BASE_URL", None): + with patch.object(base.settings, "EMBEDDINGS_DELEGATE_TO_WORKER", False): + with patch.object(base.EmbeddingsSingleton, "get_instance") as get_instance: + base.get_embeddings("some/model") + get_instance.assert_called_once() + + def test_the_delegating_client_is_shared(self): + 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") diff --git a/tests/vectorstore/test_embeddings_local.py b/tests/vectorstore/test_embeddings_local.py index b6721774..9e5a9112 100644 --- a/tests/vectorstore/test_embeddings_local.py +++ b/tests/vectorstore/test_embeddings_local.py @@ -197,7 +197,7 @@ class TestLengthSortedBatching: return wrapper, instance def test_output_order_matches_input_order(self, fake_fastembed): - with patch.object(embeddings_local.settings, "EMBEDDINGS_BATCH_SIZE", 2, create=True): + with patch.object(embeddings_local.settings, "EMBEDDINGS_MODEL_BATCH_SIZE", 2, create=True): wrapper, _ = self._wrapper(fake_fastembed, 2) texts = ["dddd", "a", "ccc", "bb", "eeeee"] out = wrapper.embed_documents(texts) @@ -206,14 +206,14 @@ class TestLengthSortedBatching: assert out == [[4.0], [1.0], [3.0], [2.0], [5.0]] def test_inputs_are_grouped_by_length_before_batching(self, fake_fastembed): - with patch.object(embeddings_local.settings, "EMBEDDINGS_BATCH_SIZE", 2, create=True): + with patch.object(embeddings_local.settings, "EMBEDDINGS_MODEL_BATCH_SIZE", 2, create=True): wrapper, instance = self._wrapper(fake_fastembed, 2) wrapper.embed_documents(["dddd", "a", "ccc", "bb", "eeeee"]) sent = instance.embed.call_args.args[0] assert [len(t) for t in sent] == [1, 2, 3, 4, 5] def test_single_batch_is_not_reordered(self, fake_fastembed): - with patch.object(embeddings_local.settings, "EMBEDDINGS_BATCH_SIZE", 32, create=True): + with patch.object(embeddings_local.settings, "EMBEDDINGS_MODEL_BATCH_SIZE", 32, create=True): wrapper, instance = self._wrapper(fake_fastembed, 32) texts = ["dddd", "a", "ccc"] out = wrapper.embed_documents(texts) @@ -221,7 +221,7 @@ class TestLengthSortedBatching: assert out == [[4.0], [1.0], [3.0]] def test_duplicate_texts_are_handled(self, fake_fastembed): - with patch.object(embeddings_local.settings, "EMBEDDINGS_BATCH_SIZE", 2, create=True): + with patch.object(embeddings_local.settings, "EMBEDDINGS_MODEL_BATCH_SIZE", 2, create=True): wrapper, _ = self._wrapper(fake_fastembed, 2) out = wrapper.embed_documents(["aa", "b", "aa", "ccc"]) assert out == [[2.0], [1.0], [2.0], [3.0]] diff --git a/tests/vectorstore/test_pgvector_live_schema.py b/tests/vectorstore/test_pgvector_live_schema.py index 273cea3e..e7192cfb 100644 --- a/tests/vectorstore/test_pgvector_live_schema.py +++ b/tests/vectorstore/test_pgvector_live_schema.py @@ -103,9 +103,20 @@ def live_dsn(postgresql, monkeypatch): @pytest.fixture def stub_embeddings(): + """Stand in for the configured model everywhere its width is read. + + ``ensure_vector_schema`` takes the width from the registry rather than by + constructing the model, so patching only the constructors would leave the + boot hook sizing the table from whatever EMBEDDINGS_NAME happens to be. + """ stub = _StubEmbeddings() with patch( "application.vectorstore.base.get_embeddings", return_value=stub + ), patch( + "application.vectorstore.base.build_local_embeddings", return_value=stub + ), patch( + "application.vectorstore.model_registry.dimension_for", + return_value=STUB_DIM, ), patch( "application.vectorstore.base.BaseVectorStore._get_embeddings", return_value=stub, @@ -140,6 +151,9 @@ class TestBootHookCreatesTheSchema: wide = _WideStubEmbeddings() with patch( "application.vectorstore.base.get_embeddings", return_value=wide + ), patch( + "application.vectorstore.model_registry.dimension_for", + return_value=wide.dimension, ): with pytest.raises(RuntimeError) as excinfo: ensure_vector_schema()