Merge branch 'main' into anydoc-support

Conflicts, and how each was taken:

- application/core/settings.py — ours. The renamed OCR_ENABLED /
  OCR_ATTACHMENTS_ENABLED / OCR_MIN_CHARS_PER_PAGE accept main's
  DOCLING_OCR_* spellings as AliasChoices, so nothing is dropped.
- application/Dockerfile — both. Main's install layers plus the
  INSTALL_DOCLING build arg.
- application/parser/file/constants.py — both imports.
- deployment/docker-compose.yaml — both. The INSTALL_DOCLING /
  INSTALL_TESSERACT build args on backend and worker, and main's
  -Q docsgpt,parsing,embeddings, which query embedding needs.
- tests/conftest.py — theirs. Both sides fixed the same pytest-postgresql
  9.0.0 autocommit= breakage; main's spelling is the one already on main.
- application/requirements.txt — the comments claimed different reasons
  torch is in core. Main's is the true one now: it removed
  sentence-transformers, so docling is torch's only remaining consumer.

Two things the merge broke without conflicting:

- onnxruntime. This branch moved it out of core into the docling extra;
  main meanwhile made it the runtime local embeddings execute on
  (fastembed). Git took the deletion, leaving fastembed with no pinned
  runtime in a repo that pins everything. Restored to core, and no longer
  pinned twice from the extra.
- The frontend copy of ATTACHMENT_PARSER_EXTENSIONS. The backend list is
  derived and picked up the anydoc suffixes; the hand-kept frontend mirror
  did not, so the composer would refuse files the API accepts.
  tests/parser/file/test_constants.py is what caught it.

ruff, pytest (9897 passed), frontend build and docs build all pass. The
image build is unverified: no Docker daemon on this machine.
This commit is contained in:
Alex committed 2026-09-04 16:42:48 +01:00
commit 3947c66cda
97 files changed
+8661 -1049

No files matched your search

+1 -1
View File
@@ -29,7 +29,7 @@ serves only the WSGI Flask app — it omits `/mcp` and the reconnect reader
### Celery (Task Queue)
```bash
celery -A application.app.celery worker -l INFO -Q docsgpt,parsing
celery -A application.app.celery worker -l INFO -Q docsgpt,parsing,embeddings
```
The `parsing` queue serves document parsing (the `read_document` tool / workflow
+3 -6
View File
@@ -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
+33 -2
View File
@@ -11,11 +11,42 @@ INTERNAL_KEY=<internal key for worker-to-backend authentication>
# NOVITA_API_KEY=<your-novita-api-key>
# OPEN_ROUTER_API_KEY=<your-openrouter-api-key>
# Remote Embeddings (Optional - for using a remote embeddings API instead of local SentenceTransformer)
# When set, the app will use the remote API and won't load SentenceTransformer (saves RAM)
# Embedding model. Leave it commented out and DocsGPT picks one for you: a
# fresh install is pinned to granite (multilingual, 32k context, the same 768
# dimensions as the legacy model), and an install that already has sources
# keeps the model its index was built with.
#
# Setting it here overrides that pin, so only set it deliberately. On an index
# that already has vectors, changing it without re-embedding leaves queries
# searching a different vector space than the stored vectors -- which fails
# silently, because both models are 768-dimensional. To switch, set it and then
# run:
# python -m application.scripts.reembed
# EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2
# Remote Embeddings (Optional - for using a remote embeddings API instead of
# running the model in-process). When set, the app calls the remote API and
# never loads a local model, which keeps the API and worker containers small.
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=
+33 -5
View File
@@ -56,23 +56,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
@@ -98,6 +112,20 @@ ruff check .
python -m pytest
```
On **macOS**, run the suite with `KMP_DUPLICATE_LIB_OK=TRUE`:
```bash
KMP_DUPLICATE_LIB_OK=TRUE python -m pytest
```
`faiss-cpu` and `torch` each ship their own LLVM OpenMP runtime, and loading
both into one process makes `libomp.dylib` abort the interpreter
(`OMP: Error #15`). It is a macOS-only packaging clash, not a code fault: Linux
resolves both to `libgomp`, which tolerates duplicates, so CI (`ubuntu-latest`)
and the Docker images are unaffected. Without the variable, whether the run
aborts depends on which tests happen to load faiss and torch in the same
process, so a green run on one selection and an abort on another is expected.
### Frontend changes
```bash
+20 -6
View File
@@ -17,11 +17,6 @@ RUN if [ -f /usr/bin/python3.12 ]; then \
echo "Python 3.12 not found"; exit 1; \
fi
# Download and unzip the model
RUN wget https://d3dg1063dc54p9.cloudfront.net/models/embeddings/mpnet-base-v2.zip && \
unzip mpnet-base-v2.zip -d models && \
rm mpnet-base-v2.zip
# Install Rust
RUN wget -q -O - https://sh.rustup.rs | sh -s -- -y
@@ -91,7 +86,26 @@ RUN groupadd -r appuser && \
# Copy the virtual environment and model from the builder stage
COPY --from=builder /venv /venv
COPY --from=builder /models /app/models
# Pre-fetch the embedding models into FastEmbed's cache so a fresh container
# does not download on first ingest and an air-gapped install works at all.
# Both defaults are baked: an upgraded deployment keeps using mpnet until it
# runs the re-embed script, while a new one starts on granite.
# The prefetch writes hub-layout snapshots (including tokenizer.json) here, so
# HF_HUB_CACHE has to point at the same directory: chunking loads the tokenizer
# through ``tokenizers``, which reads the hub cache and would otherwise fetch
# over the network on first ingest -- and fall back to cl100k when offline.
ENV EMBEDDINGS_CACHE_DIR=/app/models \
HF_HUB_CACHE=/app/models
# Only the modules the prefetch imports are copied first. It reaches nothing
# beyond model_registry, which is stdlib-only, so keeping the full source copy
# below this layer stops an unrelated edit from re-downloading ~780 MB of model
# artifacts on every build.
COPY __init__.py /app/application/__init__.py
COPY scripts/__init__.py scripts/prefetch_models.py /app/application/scripts/
COPY vectorstore/__init__.py vectorstore/model_registry.py /app/application/vectorstore/
ARG EMBEDDINGS_PREFETCH=""
RUN PYTHONPATH=/app /venv/bin/python -m application.scripts.prefetch_models ${EMBEDDINGS_PREFETCH}
# Copy your application code
COPY . /app/application
@@ -0,0 +1,35 @@
"""0031 token_usage cache tokens — persist the prompt-cache breakdown.
Providers report how much of each prompt was served from the prompt cache
(``cached_tokens``) and, on newer OpenAI-family models, how much was written
to it (``cache_write_tokens``). The LLM clients already parsed both and the
usage layer discarded them, so ``token_usage`` could not show a cache hit
rate. These two nullable columns carry the breakdown; NULL means the
provider reported nothing (distinct from 0). ``prompt_tokens`` keeps its
meaning as the provider's total. Idempotent both ways.
Revision ID: 0031_token_usage_cache_tokens
Revises: 0030_superseded_messages
"""
from typing import Sequence, Union
from alembic import op
revision: str = "0031_token_usage_cache_tokens"
down_revision: Union[str, None] = "0030_superseded_messages"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.execute("ALTER TABLE token_usage ADD COLUMN IF NOT EXISTS cached_tokens integer;")
op.execute(
"ALTER TABLE token_usage ADD COLUMN IF NOT EXISTS cache_write_tokens integer;"
)
def downgrade() -> None:
op.execute("ALTER TABLE token_usage DROP COLUMN IF EXISTS cache_write_tokens;")
op.execute("ALTER TABLE token_usage DROP COLUMN IF EXISTS cached_tokens;")
+44 -2
View File
@@ -39,6 +39,10 @@ from application.stt.stt_creator import STTCreator
from application.tts.tts_creator import TTSCreator
from application.upload_limits import (
copy_upload_to_path,
enforce_parseable_attachment,
is_unsupported_upload_message,
UnsupportedUploadTypeError,
unsupported_upload_message,
upload_limit_message,
UploadTooLargeError,
)
@@ -125,11 +129,27 @@ def _enforce_uploaded_audio_size_limit(file, filename: str) -> None:
enforce_audio_file_size_limit(size_bytes)
def _get_store_attachment_user_error(exc: Exception) -> str:
def _get_store_attachment_user_error(exc: Exception, filename: str | None = None) -> str:
"""Map an upload failure to a fixed client-facing message.
Every branch returns text this module composes itself. Nothing is read off
the exception, so no exception state — message or traceback — can reach a
response body (CodeQL py/stack-trace-exposure).
Args:
exc: The exception raised while processing one uploaded file.
filename: That file's name, which supplies the extension for the
unsupported-type message.
Returns:
The message to report for this failure.
"""
if isinstance(exc, AudioFileTooLargeError):
return build_stt_file_size_limit_message()
if isinstance(exc, UploadTooLargeError):
return upload_limit_message()
if isinstance(exc, UnsupportedUploadTypeError):
return unsupported_upload_message(filename)
return "Failed to process file"
@@ -214,6 +234,10 @@ class StoreAttachment(Resource):
with tempfile.TemporaryDirectory() as temp_dir:
staged_path = os.path.join(temp_dir, original_filename)
copy_upload_to_path(file, staged_path)
# Refuse here — nothing is stored or queued yet — so a
# video the picker let through never reaches the
# worker's reader, which would open it as text.
enforce_parseable_attachment(staged_path, original_filename)
with open(staged_path, "rb") as staged_stream:
staged_upload = FileStorage(
stream=staged_stream,
@@ -242,7 +266,9 @@ class StoreAttachment(Resource):
errors.append({
"upload_index": idx,
"filename": file.filename,
"error": _get_store_attachment_user_error(file_err),
"error": _get_store_attachment_user_error(
file_err, file.filename
),
})
if not tasks:
@@ -263,6 +289,22 @@ class StoreAttachment(Resource):
),
413,
)
if errors and all(
is_unsupported_upload_message(error.get("error")) for error in errors
):
# Tell the user which type was refused rather than the
# generic copy — this is the branch a phone picker that
# ignores ``accept`` lands in.
return make_response(
jsonify(
{
"success": False,
"message": errors[0]["error"],
"errors": errors,
}
),
400,
)
return make_response(
jsonify({"status": "error", "message": "No valid files to upload"}),
400,
+17 -3
View File
@@ -11,6 +11,7 @@ from application.api.answer.services.conversation_service import (
TERMINATED_RESPONSE_PLACEHOLDER,
)
from application.storage.db.base_repository import looks_like_uuid, row_to_dict
from application.storage.db.repositories.agents import AgentsRepository
from application.storage.db.repositories.attachments import AttachmentsRepository
from application.storage.db.repositories.conversations import ConversationsRepository
from application.storage.db.repositories.message_events import MessageEventsRepository
@@ -357,10 +358,20 @@ class SubmitFeedback(Resource):
description="Submit feedback for a conversation",
)
def post(self):
data = request.get_json() or {}
decoded_token = request.decoded_token
# api_key callers (widgets) carry no JWT; resolve the key to its owner.
scoped_api_key = None
if not decoded_token:
return make_response(jsonify({"success": False}), 401)
data = request.get_json()
api_key = data.get("api_key")
agent = None
if api_key:
with db_readonly() as conn:
agent = AgentsRepository(conn).find_by_key(api_key)
if not agent:
return make_response(jsonify({"success": False}), 401)
decoded_token = {"sub": agent.get("user_id")}
scoped_api_key = api_key
required_fields = ["feedback", "conversation_id", "question_index"]
missing_fields = check_required_fields(data, required_fields)
if missing_fields:
@@ -387,7 +398,10 @@ class SubmitFeedback(Resource):
with db_session() as conn:
repo = ConversationsRepository(conn)
conv = repo.get_any(data["conversation_id"], user_id)
if conv is None:
# A key may only rate conversations it created.
if conv is None or (
scoped_api_key and conv.get("api_key") != scoped_api_key
):
return make_response(
jsonify({"success": False, "message": "Not found"}), 404
)
+6
View File
@@ -800,6 +800,12 @@ class CreateWikiSource(Resource):
config={"kind": "wiki"},
directory_structure={},
tokens=0,
# Wiki pages are embedded like any other source, so record
# which model did it. Left unset, the column reads as NULL,
# which the boot mismatch check takes to mean "pre-dates the
# column, therefore the legacy model" -- and reports a
# correctly-embedded source as stale on every startup.
model=settings.EMBEDDINGS_NAME,
)
if initial_content:
WikiPagesRepository(conn).upsert(
+10
View File
@@ -35,6 +35,10 @@ from application.storage.db.bootstrap import ( # noqa: E402
ensure_database_ready,
ensure_vector_schema,
)
from application.storage.db.embeddings_pin import ( # noqa: E402
resolve_embeddings_pin,
warn_on_source_model_mismatch,
)
from application.stt.upload_limits import ( # noqa: E402
build_stt_file_size_limit_message,
should_reject_stt_request,
@@ -62,6 +66,12 @@ ensure_database_ready(
logger=logging.getLogger("application.app"),
)
# Which embedding model this installation uses is a property of its index, not of
# the release. Resolve it before the vector schema hook below, which sizes the
# table from EMBEDDINGS_NAME, and before anything embeds.
resolve_embeddings_pin(logging.getLogger("application.app"))
warn_on_source_model_mismatch(logging.getLogger("application.app"))
# Own the vector DB's schema here too, so the retrieval hot path is pure reads
# instead of re-running DDL for every source of every request.
if settings.AUTO_VECTOR_SCHEMA:
+17 -2
View File
@@ -115,9 +115,24 @@ def _trim_native_heap() -> None:
pass
# Tasks that allocate almost nothing and run on a latency-sensitive path, so
# the reclaim below costs far more than it recovers. Query embedding is one:
# measured at ~86 ms for the collect against ~8 ms for the embed itself on a
# worker holding the ONNX model, i.e. a 9x slowdown of the whole round trip.
_NO_RECLAIM_TASKS = frozenset({"application.vectorstore.embeddings_tasks.embed_texts"})
@task_postrun.connect
def _reclaim_memory_after_task(*args, **kwargs):
"""Drop per-task allocations so the prefork child's RSS doesn't ratchet."""
def _reclaim_memory_after_task(task=None, **kwargs):
"""Drop per-task allocations so the prefork child's RSS doesn't ratchet.
Skipped for the tasks in :data:`_NO_RECLAIM_TASKS`. This exists for the
large transient allocations docling/torch parsing makes; running a full
generational collect after a task that allocated a few kilobytes just
charges the next task for walking the whole heap.
"""
if getattr(task, "name", None) in _NO_RECLAIM_TASKS:
return
gc.collect()
torch = sys.modules.get("torch")
if torch is not None:
+13 -2
View File
@@ -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"
+109 -164
View File
@@ -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,41 +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
# 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)
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
# 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
@@ -94,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)
@@ -207,30 +219,22 @@ class Settings(BaseSettings):
# unaffected and always uses docling, because chunking and retrieval do
# depend on that structure.
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
@@ -293,9 +297,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
@@ -315,16 +318,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"
@@ -338,11 +336,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
@@ -371,34 +366,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",
@@ -411,81 +393,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=<oldest>`` 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
@@ -500,23 +463,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)
@@ -526,44 +485,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"
@@ -572,8 +518,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
+4 -2
View File
@@ -157,12 +157,14 @@ class GraphStore:
"""Dimension of the configured embeddings model, matching ``PGVectorStore``.
Falls back to ``DEFAULT_NAME_EMBEDDING_DIM`` so the graph table and the
pgvector ``documents`` table always agree on the configured model.
pgvector ``documents`` table always agree on the configured model. A
model outside the registry reports ``None`` rather than no attribute,
so the fallback cannot be left to ``getattr``.
"""
from application.vectorstore.base import get_embeddings
embedding = get_embeddings()
return getattr(embedding, "dimension", DEFAULT_NAME_EMBEDDING_DIM)
return getattr(embedding, "dimension", None) or DEFAULT_NAME_EMBEDDING_DIM
@staticmethod
def create_schema(conn, *, dimension: int = DEFAULT_NAME_EMBEDDING_DIM) -> None:
+10 -6
View File
@@ -6,7 +6,7 @@ from typing import Any, Dict, Generator, List, Optional, Tuple
from anthropic import Anthropic
from application.core.settings import settings
from application.llm.base import BaseLLM
from application.llm.base import BaseLLM, optional_int
from application.storage.storage_creator import StorageCreator
logger = logging.getLogger(__name__)
@@ -387,8 +387,12 @@ class AnthropicLLM(BaseLLM):
# message_start are reused below instead.
if not output_only:
base_input = int(getattr(usage, "input_tokens", 0) or 0)
created = int(getattr(usage, "cache_creation_input_tokens", 0) or 0)
cached = int(getattr(usage, "cache_read_input_tokens", 0) or 0)
# None (unreported) and 0 (reported, no caching) must stay
# distinguishable — see ``optional_int``.
created = optional_int(
getattr(usage, "cache_creation_input_tokens", None)
)
cached = optional_int(getattr(usage, "cache_read_input_tokens", None))
except (TypeError, ValueError):
return
@@ -398,11 +402,11 @@ class AnthropicLLM(BaseLLM):
prompt = int(previous.get("prompt_tokens", 0) or 0)
details = previous.get("prompt_tokens_details")
else:
prompt = base_input + created + cached
prompt = base_input + (created or 0) + (cached or 0)
details = {}
if cached:
if cached is not None:
details["cached_tokens"] = cached
if created:
if created is not None:
details["cache_creation_tokens"] = created
details = details or None
+29
View File
@@ -33,6 +33,29 @@ _STREAM_RETRYABLE_TRANSPORT_ERRORS = (
)
def optional_int(value) -> Optional[int]:
"""Coerce a provider-reported count to int, keeping "unreported" as None.
Cache bins persist as NULL when a provider says nothing and as 0 when it
reports zero, so the two must stay distinguishable all the way from the
usage object: providers report ``cached_tokens: 0`` on every uncached
request, and folding that into "unknown" would file the ordinary case as
no-data.
Args:
value: The raw attribute off a provider usage object.
Returns:
The integer value, or None when absent or not a number.
"""
if value is None:
return None
try:
return int(value)
except (TypeError, ValueError):
return None
class BaseLLM(ABC):
# Stamped onto the ``llm_stream_start`` event so dashboards can group
# calls by vendor. Subclasses override.
@@ -710,6 +733,7 @@ class BaseLLM(ABC):
completion_tokens,
latency_ms,
cached_tokens=None,
cache_write_tokens=None,
error=None,
):
# Non-streaming counterpart to ``_emit_stream_finished_log``. Paired
@@ -730,6 +754,8 @@ class BaseLLM(ABC):
}
if cached_tokens is not None:
extra["cached_tokens"] = int(cached_tokens)
if cache_write_tokens is not None:
extra["cache_write_tokens"] = int(cache_write_tokens)
if error is not None:
extra["error_class"] = type(error).__name__
logging.info("llm_gen_finished", extra=extra)
@@ -757,6 +783,7 @@ class BaseLLM(ABC):
completion_tokens,
latency_ms,
cached_tokens=None,
cache_write_tokens=None,
error=None,
):
# Paired with ``llm_stream_start`` so cost dashboards can sum tokens
@@ -774,6 +801,8 @@ class BaseLLM(ABC):
}
if cached_tokens is not None:
extra["cached_tokens"] = int(cached_tokens)
if cache_write_tokens is not None:
extra["cache_write_tokens"] = int(cache_write_tokens)
if error is not None:
extra["error_class"] = type(error).__name__
logging.info("llm_stream_finished", extra=extra)
+38 -13
View File
@@ -9,7 +9,7 @@ from typing import Any, Callable
from openai import BadRequestError, OpenAI
from application.core.settings import settings
from application.llm.base import BaseLLM
from application.llm.base import BaseLLM, optional_int
from application.storage.storage_creator import StorageCreator
# Placeholder sent to OpenAI-compatible backends that require no credentials.
@@ -1305,19 +1305,38 @@ class OpenAILLM(BaseLLM):
}
input_details = getattr(usage, "prompt_tokens_details", None)
output_details = getattr(usage, "completion_tokens_details", None)
try:
cached = int(getattr(input_details, "cached_tokens", 0) or 0)
reasoning = int(getattr(output_details, "reasoning_tokens", 0) or 0)
except (TypeError, ValueError):
cached = 0
reasoning = 0
if cached:
result["prompt_tokens_details"] = {"cached_tokens": cached}
cached = optional_int(getattr(input_details, "cached_tokens", None))
written = optional_int(getattr(input_details, "cache_write_tokens", None))
reasoning = optional_int(
getattr(output_details, "reasoning_tokens", None)
)
details = self._prompt_cache_details(cached, written)
if details:
result["prompt_tokens_details"] = details
if reasoning:
result["completion_tokens_details"] = {"reasoning_tokens": reasoning}
self._last_usage = result
self._last_usage_claimed = False
@staticmethod
def _prompt_cache_details(cached: int | None, written: int | None) -> dict:
"""Build the ``prompt_tokens_details`` breakdown from the two cache bins.
A bin is included whenever the provider reported it, zero included:
downstream, an absent bin persists as NULL ("we don't know") and a
reported zero as 0 ("no cache hits"), and OpenAI reports
``cached_tokens: 0`` on every uncached request. Collapsing the two
would file every ordinary request as unknown. ``cache_write_tokens``
is reported by newer OpenAI-family deployments that charge for cache
writes.
"""
details = {}
if cached is not None:
details["cached_tokens"] = cached
if written is not None:
details["cache_write_tokens"] = written
return details
@staticmethod
def _function_call_ids(response):
"""call_ids of every ``function_call`` item in a Responses output.
@@ -1356,10 +1375,16 @@ class OpenAILLM(BaseLLM):
"completion_tokens": completion,
"total_tokens": int(getattr(usage, "total_tokens", 0) or prompt + completion),
}
cached = int(getattr(input_details, "cached_tokens", 0) or 0)
reasoning = int(getattr(output_details, "reasoning_tokens", 0) or 0)
if cached:
result["prompt_tokens_details"] = {"cached_tokens": cached}
cached = optional_int(getattr(input_details, "cached_tokens", None))
written = optional_int(
getattr(input_details, "cache_write_tokens", None)
)
reasoning = optional_int(
getattr(output_details, "reasoning_tokens", None)
)
details = self._prompt_cache_details(cached, written)
if details:
result["prompt_tokens_details"] = details
if reasoning:
result["completion_tokens_details"] = {"reasoning_tokens": reasoning}
self._last_usage = result
+81 -25
View File
@@ -3,10 +3,17 @@ from typing import List, Tuple
import logging
from application.parser.chunking_creator import ChunkerCreator
from application.parser.schema.base import Document
from application.utils import get_encoding
from application.parser.tokenization import get_token_counter
logger = logging.getLogger(__name__)
# Smallest share of ``max_tokens`` a chunk must keep for body text when the
# header is duplicated onto every chunk. A header that leaves less than this
# makes each chunk mostly repeated header and multiplies the chunk count -- at
# a budget of 32 out of 1250 a document splits into 39x more chunks than it
# needs -- so duplication is dropped rather than honoured.
_MIN_BODY_BUDGET_RATIO = 0.25
class Chunker:
"""Classic token-window chunker (registered as ``classic_chunk``).
@@ -24,10 +31,15 @@ class Chunker:
duplicate_headers: bool = False,
):
self.chunking_strategy = chunking_strategy
self.max_tokens = max_tokens
# A budget below 1 would ask for a chunk per token; the strategy
# chunkers clamp the same way.
self.max_tokens = max(1, int(max_tokens))
self.min_tokens = min_tokens
self.duplicate_headers = duplicate_headers
self.encoding = get_encoding()
# Counted in the embedding model's tokenizer, not cl100k: ``max_tokens``
# is compared against a limit the embedding server enforces in its own
# units, so counting in any other unit is a guess.
self.counter = get_token_counter()
def separate_header_and_body(self, text: str) -> Tuple[str, str]:
header_pattern = r"^(.*?\n){3}"
@@ -42,28 +54,73 @@ class Chunker:
def split_document(self, doc: Document) -> List[Document]:
split_docs = []
header, body = self.separate_header_and_body(doc.text)
header_tokens = self.encoding.encode_ordinary(header) if header else []
body_tokens = self.encoding.encode_ordinary(body)
"""Split one oversized document into ``max_tokens``-sized chunks.
current_position = 0
part_index = 0
while current_position < len(body_tokens):
end_position = current_position + self.max_tokens - len(header_tokens)
chunk_tokens = (header_tokens + body_tokens[current_position:end_position]
if self.duplicate_headers or part_index == 0 else body_tokens[current_position:end_position])
chunk_text = self.encoding.decode(chunk_tokens)
new_doc = Document(
text=chunk_text,
doc_id=f"{doc.doc_id}-{part_index}",
embedding=doc.embedding,
extra_info={**(doc.extra_info or {}), "token_count": len(chunk_tokens)}
Pieces are sliced out of the original text rather than decoded back
from token ids. WordPiece tokenizers normalise as they decode --
all-mpnet-base-v2 lowercases -- so a decode round-trip would rewrite
every stored document.
"""
header, body = self.separate_header_and_body(doc.text)
header_tokens = self.counter.count(header) if header else 0
if header and header_tokens >= self.max_tokens:
# The header alone fills the budget, so no cut of the body can keep
# a chunk within it and duplicating it would leave a one-token body
# budget -- a chunk per body token. It is only the first three
# lines, not something worth preserving at that cost, so it goes
# back to being ordinary text and the document splits evenly.
logger.warning(
"Header of %s is %d token(s), at or over the %d-token chunk "
"budget; treating it as body text.",
doc.doc_id,
header_tokens,
self.max_tokens,
)
body = f"{header}{body}"
header, header_tokens = "", 0
# A chunk carrying the header has that much less room for body text.
with_header_budget = max(1, self.max_tokens - header_tokens)
duplicate_headers = self.duplicate_headers
if duplicate_headers and with_header_budget < self.max_tokens * _MIN_BODY_BUDGET_RATIO:
logger.warning(
"Header of %s leaves only %d of %d tokens for body text; "
"carrying it on the first chunk only.",
doc.doc_id,
with_header_budget,
self.max_tokens,
)
duplicate_headers = False
if duplicate_headers:
body_pieces = self.counter.split(body, with_header_budget)
else:
body_pieces = self.counter.split(
body, self.max_tokens, first_max_tokens=with_header_budget
)
if not body_pieces and header:
# Nothing but a header: the loop below only ever emits the header
# attached to a body piece, so without this the document is dropped
# from the index entirely.
body_pieces = [""]
split_docs = []
for part_index, piece in enumerate(body_pieces):
include_header = bool(header) and (duplicate_headers or part_index == 0)
chunk_text = f"{header}{piece}" if include_header else piece
split_docs.append(
Document(
text=chunk_text,
doc_id=f"{doc.doc_id}-{part_index}",
embedding=doc.embedding,
extra_info={
**(doc.extra_info or {}),
"token_count": self.counter.count(chunk_text),
},
)
)
split_docs.append(new_doc)
current_position = end_position
part_index += 1
header_tokens = []
return split_docs
def classic_chunk(self, documents: List[Document]) -> List[Document]:
@@ -71,8 +128,7 @@ class Chunker:
i = 0
while i < len(documents):
doc = documents[i]
tokens = self.encoding.encode_ordinary(doc.text)
token_count = len(tokens)
token_count = self.counter.count(doc.text)
if self.min_tokens <= token_count <= self.max_tokens:
doc.extra_info = doc.extra_info or {}
+9 -17
View File
@@ -16,7 +16,7 @@ from typing import List
from application.parser.chunking import Chunker
from application.parser.chunking_creator import ChunkerCreator
from application.parser.schema.base import Document
from application.utils import get_encoding
from application.parser.tokenization import get_token_counter
logger = logging.getLogger(__name__)
@@ -41,19 +41,16 @@ class _BaseStrategyChunker:
self.max_tokens = max(1, int(max_tokens))
self.min_tokens = max(0, int(min_tokens))
self.duplicate_headers = duplicate_headers
self.encoding = get_encoding()
# Same unit as the embedding server counts in; see
# ``application.parser.tokenization``.
self.counter = get_token_counter()
def _token_count(self, text: str) -> int:
return len(self.encoding.encode_ordinary(text))
return self.counter.count(text)
def _split_by_tokens(self, text: str) -> List[str]:
"""Split ``text`` into pieces no larger than ``max_tokens`` tokens."""
tokens = self.encoding.encode_ordinary(text)
pieces = []
for start in range(0, len(tokens), self.max_tokens):
chunk_tokens = tokens[start:start + self.max_tokens]
pieces.append(self.encoding.decode(chunk_tokens))
return pieces
return self.counter.split(text, self.max_tokens)
def _emit(self, base: Document, part_index: int, text: str) -> Document:
"""Build a child Document carrying token_count and inherited info."""
@@ -188,16 +185,11 @@ class ParentChildChunker(_BaseStrategyChunker):
processed: List[Document] = []
child_size = self._child_size()
for doc in documents:
tokens = self.encoding.encode_ordinary(doc.text)
part_index = 0
for p_start in range(0, len(tokens), self.max_tokens):
parent_tokens = tokens[p_start:p_start + self.max_tokens]
parent_text = self.encoding.decode(parent_tokens)
for parent_text in self.counter.split(doc.text, self.max_tokens):
if not parent_text.strip():
continue
for c_start in range(0, len(parent_tokens), child_size):
child_tokens = parent_tokens[c_start:c_start + child_size]
child_text = self.encoding.decode(child_tokens)
for child_text in self.counter.split(parent_text, child_size):
if not child_text.strip():
continue
child = Document(
@@ -208,7 +200,7 @@ class ParentChildChunker(_BaseStrategyChunker):
embedding=doc.embedding,
extra_info={
**(doc.extra_info or {}),
"token_count": len(child_tokens),
"token_count": self.counter.count(child_text),
"parent_text": parent_text,
},
)
+68
View File
@@ -1,5 +1,7 @@
"""Shared file-extension constants for parsing and ingestion flows."""
import os
from application.parser.file.anydoc_parser import ANYDOC_GAINED_SUFFIXES
from application.stt.constants import SUPPORTED_AUDIO_EXTENSIONS
@@ -32,3 +34,69 @@ SUPPORTED_SOURCE_EXTENSIONS = (
*SUPPORTED_SOURCE_IMAGE_EXTENSIONS,
*SUPPORTED_AUDIO_EXTENSIONS,
)
# Suffixes the attachment path has a dedicated parser for — exactly the keys
# of ``get_default_file_extractor()``. Kept as a literal (importing ``bulk``
# here would drag docling into the API process), with
# ``tests/parser/file/test_constants.py`` asserting the two agree.
#
# This is *not* the whole attachment allow-list: a suffix with no parser is
# read by ``SimpleDirectoryReader``'s plain-text fallthrough, which is right
# for a .py or a .log and catastrophic for a video — the reader "extracts"
# megabytes of binary garbage, truncates it, and stores it with
# ``extraction.status == "ok"``. So unparsed suffixes are admitted on
# content instead (``upload_limits.enforce_parseable_attachment``): text
# passes, binary is refused. Zip is deliberately absent — source ingestion
# extracts archives, the attachment path does not, and a zip fails the
# content check like any other binary.
#
# Mirrored in ``frontend/src/constants/fileUpload.ts``; update both together.
ATTACHMENT_PARSER_EXTENSIONS = frozenset(
{
*SUPPORTED_SOURCE_EXTENSIONS,
".xhtml",
".adoc",
".asciidoc",
".tiff",
".tif",
".bmp",
".webp",
".vtt",
".xml",
}
# .txt has no parser of its own — it *is* the plain-text fallthrough. It
# must be sniffed like any other unparsed suffix, or renaming a video to
# notes.txt walks straight back into the bug this gate exists for.
- {".txt"}
)
def attachment_extension(filename: str | None) -> str:
"""Return the lower-cased extension of ``filename`` including the dot, or ``""``.
Args:
filename: A bare filename or path; ``None`` and empty strings yield ``""``.
Returns:
The last suffix in lower case (``".pdf"``), or ``""`` when there is none.
"""
if not filename:
return ""
return os.path.splitext(os.path.basename(str(filename)))[1].lower()
def has_attachment_parser(filename: str | None) -> bool:
"""Return whether an attachment's suffix has a dedicated parser.
A False result does not mean the file is refused: it means nothing but the
plain-text fallthrough will read it, so it has to earn its place on
content. The decision is by extension only, never by the mime type a
browser reports (mobile pickers ignore ``accept`` and lie).
Args:
filename: The upload's filename.
Returns:
True when the suffix is in ``ATTACHMENT_PARSER_EXTENSIONS``.
"""
return attachment_extension(filename) in ATTACHMENT_PARSER_EXTENSIONS
+275
View File
@@ -0,0 +1,275 @@
"""Counting and splitting text in the embedding model's own tokenizer.
Chunk sizes only mean something in the units the embedding server counts. The
chunker used to count cl100k (tiktoken) while the server counted whatever its
model used, so ``max_tokens`` was a value in one unit compared against a limit
in another. Every recalibration of that number was really an attempt to guess
the conversion factor.
This module removes the conversion. :func:`get_token_counter` returns a counter
backed by the configured embedding model's tokenizer, falling back to cl100k
when that tokenizer cannot be loaded -- an offline install, or a model the
registry does not describe.
Splitting never round-trips through ``decode``. Byte-level BPE decodes
losslessly, but WordPiece tokenizers do not: all-mpnet-base-v2 lowercases, so
decoding ``"Hello World"`` yields ``"hello world"`` and would silently rewrite
every stored document. Counters therefore cut the *original* string at
character offsets the tokenizer reports.
"""
from __future__ import annotations
import logging
import threading
from typing import Iterator, List, Optional, Tuple
from application.core.settings import settings
from application.utils import get_encoding
from application.vectorstore.model_registry import resolve
logger = logging.getLogger(__name__)
_cache: dict = {}
_cache_lock = threading.Lock()
def _windows(total: int, first: int, rest: int) -> Iterator[Tuple[int, int]]:
"""Yield ``(start, end)`` token windows, the first sized independently.
A header consumes part of the first chunk's budget but none of the rest,
so the first window is often smaller than those that follow.
"""
start = 0
budget = max(1, first)
while start < total:
end = min(start + budget, total)
yield start, end
start = end
budget = max(1, rest)
# A token spanning more characters than this collapsed a run the tokenizer
# could not break up: WordPiece emits a single ``[UNK]`` for any word longer
# than ``max_input_chars_per_word``. Charging that once makes a base64 blob or
# a minified bundle look tiny, so nothing splits it and an oversized chunk
# reaches the embedding server. Real tokens are a few characters, so prose
# never reaches this bound.
_MAX_CHARS_PER_TOKEN = 16
def _token_weight(start: int, end: int) -> int:
"""Tokens a span costs, charging a collapsed run by its length."""
span = end - start
return max(1, -(-span // _MAX_CHARS_PER_TOKEN))
def _cap_piece_chars(pieces: List[str], first: int, rest: int) -> List[str]:
"""Cut any piece holding more characters than its budget can cover."""
capped: List[str] = []
budget = first
for piece in pieces:
limit = max(1, budget * _MAX_CHARS_PER_TOKEN)
while len(piece) > limit:
capped.append(piece[:limit])
piece = piece[limit:]
budget = rest
limit = max(1, budget * _MAX_CHARS_PER_TOKEN)
capped.append(piece)
budget = rest
return capped
class TokenCounter:
"""Counts and splits text in one tokenizer's units.
Attributes:
name: Human-readable identifier of the underlying tokenizer, for logs.
"""
name = "unknown"
def count(self, text: str) -> int:
"""Number of tokens ``text`` occupies."""
raise NotImplementedError
def split(
self, text: str, max_tokens: int, first_max_tokens: Optional[int] = None
) -> List[str]:
"""Cut ``text`` into consecutive pieces of at most ``max_tokens``.
The concatenation of the returned pieces equals ``text`` exactly.
Args:
text: Source text.
max_tokens: Token budget per piece; values below 1 are treated as 1.
first_max_tokens: Budget for the first piece only, when it must
leave room for a header. Defaults to ``max_tokens``.
Returns:
The pieces, in order. An empty ``text`` yields an empty list.
"""
raise NotImplementedError
class TiktokenCounter(TokenCounter):
"""cl100k counter -- the historical behaviour, and the fallback."""
name = "cl100k_base"
def __init__(self) -> None:
self._encoding = get_encoding()
def count(self, text: str) -> int:
if not text:
return 0
return len(self._encoding.encode_ordinary(text))
def split(
self, text: str, max_tokens: int, first_max_tokens: Optional[int] = None
) -> List[str]:
if not text:
return []
rest = max(1, max_tokens)
first = max(1, first_max_tokens if first_max_tokens is not None else rest)
tokens = self._encoding.encode_ordinary(text)
if len(tokens) <= first:
return [text]
# Cut the original string at the character offsets the tokenizer
# reports, the way :class:`HuggingFaceCounter` does. Decoding each
# window on its own instead splits any multi-byte character that
# straddles a boundary across two byte sequences, and each half decodes
# to U+FFFD -- roughly one boundary in five on CJK text, silently
# destroying a character per cut.
offsets = self._encoding.decode_with_offsets(tokens)[1]
end_of_text = len(text)
pieces: List[str] = []
cursor = 0
for _, end in _windows(len(tokens), first, rest):
# Offsets index into the decoded string. ``encode_ordinary``
# round-trips for anything cl100k can represent, so it is ``text``;
# clamping keeps the cuts in range if it ever is not, and slicing
# ``text`` monotonically keeps the pieces reassembling exactly
# either way.
end_char = end_of_text if end >= len(offsets) else min(offsets[end], end_of_text)
if end_char <= cursor:
continue
pieces.append(text[cursor:end_char])
cursor = end_char
if cursor < end_of_text:
if pieces:
pieces[-1] = pieces[-1] + text[cursor:]
else:
pieces.append(text[cursor:])
return pieces
class HuggingFaceCounter(TokenCounter):
"""Counts in a Hugging Face tokenizer, slicing by character offsets."""
def __init__(self, tokenizer, name: str) -> None:
self._tokenizer = tokenizer
self.name = name
def _encode(self, text: str):
return self._tokenizer.encode(text, add_special_tokens=False)
def count(self, text: str) -> int:
if not text:
return 0
encoding = self._encode(text)
if not encoding.offsets:
return len(encoding.ids)
return sum(_token_weight(start, end) for start, end in encoding.offsets)
def split(
self, text: str, max_tokens: int, first_max_tokens: Optional[int] = None
) -> List[str]:
if not text:
return []
rest = max(1, max_tokens)
first = max(1, first_max_tokens if first_max_tokens is not None else rest)
offsets = self._encode(text).offsets
if self.count(text) <= first:
return [text]
pieces: List[str] = []
cursor = 0
for start_token, end_token in _windows(len(offsets), first, rest):
window = offsets[start_token:end_token]
# Some tokenizers emit (0, 0) for specials or normalised-away
# characters; those carry no span to cut on.
spans = [end for _, end in window if end > cursor]
if not spans:
continue
end_char = max(spans)
pieces.append(text[cursor:end_char])
cursor = end_char
if cursor < len(text):
# Trailing characters the tokenizer dropped (e.g. whitespace) belong
# to the final piece, so no input is lost.
if pieces:
pieces[-1] = pieces[-1] + text[cursor:]
else:
pieces.append(text[cursor:])
return _cap_piece_chars(pieces, first, rest)
def _load_hf_counter(repo: str) -> Optional[HuggingFaceCounter]:
"""Load ``repo``'s tokenizer, or ``None`` if it is not reachable."""
try:
from tokenizers import Tokenizer
tokenizer = Tokenizer.from_pretrained(repo)
# Repos ship padding and truncation defaults meant for inference
# batches. Left on, every count returns the padded width (128 for
# mpnet) and the offsets carry (0, 0) entries for the padding, so both
# counting and slicing are wrong.
tokenizer.no_padding()
tokenizer.no_truncation()
return HuggingFaceCounter(tokenizer, repo)
except Exception as exc: # noqa: BLE001 -- chunking must never hard-fail here
logger.warning(
"Could not load the tokenizer for %s (%s); counting chunk sizes in "
"cl100k instead. Chunk sizes will be approximate for this model.",
repo,
exc,
)
return None
def get_token_counter(embeddings_name: Optional[str] = None) -> TokenCounter:
"""Return the counter for ``embeddings_name``, cached per process.
Args:
embeddings_name: Model name; defaults to ``settings.EMBEDDINGS_NAME``.
Returns:
A :class:`HuggingFaceCounter` for a model whose tokenizer could be
loaded, else a :class:`TiktokenCounter`.
"""
name = embeddings_name or getattr(settings, "EMBEDDINGS_NAME", None)
key = name or "__default__"
with _cache_lock:
if key in _cache:
return _cache[key]
spec = resolve(name)
repo = spec.repo if spec else name
counter: TokenCounter
if repo and (spec is None or spec.provider == "fastembed"):
counter = _load_hf_counter(repo) or TiktokenCounter()
else:
counter = TiktokenCounter()
with _cache_lock:
_cache[key] = counter
logger.info("Chunking will count tokens with %s", counter.name)
return counter
def reset_cache() -> None:
"""Drop cached counters. For tests and for a settings change at runtime."""
with _cache_lock:
_cache.clear()
+2 -3
View File
@@ -7,10 +7,9 @@
# not either: OCR_ENABLED=true with the tesseract binary (or a DeepSeek-OCR
# endpoint) runs through application/parser/file/ocr_parser.py.
#
# Install ON TOP of requirements.txt (torch/transformers stay pinned there,
# shared with sentence-transformers):
# Install ON TOP of requirements.txt, which pins torch, transformers and
# onnxruntime (rapidocr runs on the same onnxruntime fastembed uses):
# pip install -r requirements.txt -r requirements-docling.txt
# Docker: build with --build-arg INSTALL_DOCLING=true
docling==2.119.0
rapidocr==3.9.2
onnxruntime==1.28.0
+10 -7
View File
@@ -21,6 +21,10 @@ ddgs>=8.0.0
fast-ebook
elevenlabs==2.62.0
Flask==3.1.3
fastembed==0.8.0
# fastembed's runtime: local embeddings execute on it, so it stays pinned in
# core even though the docling extra's rapidocr runs on the same package.
onnxruntime==1.28.0
faiss-cpu==1.15.0
fastmcp==3.4.6
flask-restx==1.3.2
@@ -91,19 +95,18 @@ referencing>=0.28.0,<0.38.0
regex==2026.7.19
requests==2.34.2
retry==0.9.2
sentence-transformers==5.7.0
sqlalchemy>=2.0,<3
starlette>=1.0,<2
tiktoken==0.13.0
tokenizers==0.22.2
# torch stays in core: sentence-transformers (local embeddings) needs it even
# with the docling extra uninstalled.
# torch and transformers serve docling's model stack alone -- embeddings moved
# to fastembed/onnxruntime. Still in core rather than requirements-docling.txt,
# so installing the optional extra never re-resolves the shared pins; both can
# move out with docling.
torch==2.11.0
tqdm==4.67.3
# Kept in core for sentence-transformers, so installing the optional docling
# extra never changes the shared model stack. Capped <5.9.0 to match
# docling-core's own darwin pin: 5.9+ breaks docling's PDF layout model on
# Apple Silicon ("Cannot convert a MPS Tensor to float64").
# Capped <5.9.0 to match docling-core's own darwin pin: 5.9+ breaks docling's
# PDF layout model on Apple Silicon ("Cannot convert a MPS Tensor to float64").
transformers==5.8.1
typing-extensions==4.16.0
typing-inspect==0.9.0
View File
Whitespace-only changes.
+89
View File
@@ -0,0 +1,89 @@
"""Download embedding model artifacts into FastEmbed's cache.
Run at image build time so a fresh container does not download a model on its
first ingest, and an air-gapped install works at all. Both the legacy and the
current default are baked: an upgraded deployment keeps using mpnet until it
runs ``reembed``, while a new one starts on granite.
Usage::
python -m application.scripts.prefetch_models # the defaults
python -m application.scripts.prefetch_models granite-311m # a subset
"""
from __future__ import annotations
import logging
import sys
from typing import List, Optional, Sequence
from application.vectorstore.model_registry import (
DEFAULT_LEGACY,
DEFAULT_NEW_INSTALL,
known_names,
resolve,
)
logger = logging.getLogger("prefetch_models")
#: Fetched when no names are given.
DEFAULT_MODELS = (DEFAULT_LEGACY, DEFAULT_NEW_INSTALL)
def prefetch(names: Sequence[str], cache_dir: Optional[str] = None) -> List[str]:
"""Fetch each named model's artifacts.
Args:
names: Registry names or aliases.
cache_dir: FastEmbed cache directory; its default when omitted.
Returns:
The repositories actually fetched.
Raises:
SystemExit: If a name is not in the registry, since a silent skip at
build time becomes a download at run time on an offline host.
"""
from fastembed import TextEmbedding
from fastembed.common.model_description import ModelSource, PoolingType
pooling_types = {"cls": PoolingType.CLS, "mean": PoolingType.MEAN}
fetched: List[str] = []
for name in names:
spec = resolve(name)
if spec is None:
raise SystemExit(
f"Unknown embedding model {name!r}. Known: {', '.join(known_names())}"
)
if spec.provider != "fastembed":
logger.info("Skipping %s: served remotely, nothing to cache.", spec.name)
continue
logger.info("Fetching %s", spec.repo)
TextEmbedding.add_custom_model(
model=spec.repo,
pooling=pooling_types[spec.pooling],
normalization=spec.normalize,
sources=ModelSource(hf=spec.repo),
dim=spec.dimension,
model_file=spec.onnx_file,
)
kwargs = {"model_name": spec.repo}
if cache_dir:
kwargs["cache_dir"] = cache_dir
TextEmbedding(**kwargs)
fetched.append(spec.repo)
return fetched
def main(argv: Optional[Sequence[str]] = None) -> int:
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
import os
names = list(argv) if argv else list(DEFAULT_MODELS)
fetched = prefetch(names, os.environ.get("EMBEDDINGS_CACHE_DIR"))
logger.info("Cached %d model(s): %s", len(fetched), ", ".join(fetched))
return 0
if __name__ == "__main__":
sys.exit(main(sys.argv[1:]))
+542
View File
@@ -0,0 +1,542 @@
"""Re-embed an existing index with the currently configured embedding model.
Changing ``EMBEDDINGS_NAME`` does not change vectors already stored. Because
several models share a width -- granite-311m and all-mpnet-base-v2 are both
768-dimensional -- a swapped model does not fail any dimension check. It simply
retrieves against vectors that mean something else, which shows up as bad
answers rather than as an error.
This script rebuilds the vectors from the chunk text already held in the store,
so it never re-downloads, re-parses or re-chunks anything: no source files, no
crawl, no docling. It works against pgvector and FAISS. With ``GRAPHRAG_ENABLED``
it also rewrites ``graph_nodes.name_embedding``, which seeds graph traversal and
would otherwise be left behind in the previous model's space.
Usage::
python -m application.scripts.reembed --dry-run # report, change nothing
python -m application.scripts.reembed # re-embed everything
python -m application.scripts.reembed --sources a,b # only these sources
"""
from __future__ import annotations
import argparse
import logging
import sys
import time
from dataclasses import dataclass
from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple
from psycopg import sql
from application.core.settings import settings
from application.vectorstore.model_registry import resolve
from application.vectorstore.vector_creator import VectorCreator
logger = logging.getLogger("reembed")
#: Stores this script knows how to rewrite in place.
SUPPORTED_STORES = ("pgvector", "faiss")
#: Chunks embedded (and, for pgvector, written) per transaction.
DEFAULT_BATCH_SIZE = 64
@dataclass
class _RebuildDoc:
"""Minimal document shape ``FaissStore(docs_init=...)`` accepts."""
page_content: str
metadata: Dict[str, Any]
class ReembedError(RuntimeError):
"""Raised for a misconfiguration the user needs to fix before running."""
def _log_setup(verbose: bool) -> None:
logging.basicConfig(
level=logging.DEBUG if verbose else logging.INFO,
format="%(asctime)s %(levelname)s %(message)s",
)
def _batched(items: Sequence[Any], size: int) -> Iterable[Sequence[Any]]:
for start in range(0, len(items), size):
yield items[start : start + size]
def list_source_ids(store_type: str) -> List[str]:
"""Every source id that has vectors in the configured store.
Args:
store_type: ``"pgvector"`` or ``"faiss"``.
Returns:
Source ids, sorted for a stable and resumable ordering.
"""
if store_type == "pgvector":
return _pgvector_source_ids()
return _faiss_source_ids()
def _pgvector_source_ids() -> List[str]:
store = VectorCreator.create_vectorstore("pgvector", source_id="")
# ``_get_connection`` hands back a pooled connection, not a context
# manager: closing it via ``with`` would cost the pool a slot.
conn = store._get_connection()
cursor = conn.cursor()
try:
# Table and column names cannot be bound as parameters, so they go
# through ``sql.Identifier``, which quotes them. They are internal
# constants rather than user input, but composing SQL by f-string is
# the habit worth not having.
cursor.execute(
sql.SQL(
"SELECT DISTINCT source_id FROM {table} "
"WHERE source_id IS NOT NULL ORDER BY source_id"
).format(table=sql.Identifier(store._table_name))
)
return [row[0] for row in cursor.fetchall()]
finally:
cursor.close()
store.close()
def _faiss_source_ids() -> List[str]:
"""Directories under the FAISS root that hold an index."""
from application.storage.storage_creator import StorageCreator
storage = StorageCreator.get_storage()
root = "indexes"
ids: List[str] = []
try:
for path in storage.list_files(root):
# ``indexes/<source_id>/index.faiss``
parts = str(path).split("/")
if parts and parts[-1].endswith(".faiss") and len(parts) >= 2:
ids.append(parts[-2])
except Exception as exc: # noqa: BLE001 -- surfaced to the user below
raise ReembedError(
f"Could not list FAISS indexes under {root!r}: {exc}"
) from exc
return sorted(set(ids))
def _graph_nodes_exist(conn) -> bool:
"""True when the GraphRAG node table is present in this database."""
cursor = conn.cursor()
try:
cursor.execute("SELECT to_regclass('public.graph_nodes')")
row = cursor.fetchone()
return bool(row and row[0])
finally:
cursor.close()
def reembed_graph_nodes(store, conn, source_id: str, batch_size: int, dry_run: bool) -> int:
"""Rewrite one source's GraphRAG node name vectors.
``graph_nodes.name_embedding`` is what every graph traversal starts from:
the retriever embeds the query and takes the nearest node names as seeds.
It is written once at extraction time and never revisited, so re-embedding
only the chunk table leaves the graph seeding against the previous model.
Nothing errors -- mpnet and granite-311m are both 768-dimensional, so the
column accepts the mismatch -- the graph just walks from the wrong nodes.
The embedded text is ``graph_nodes.name``, which is already stored, so this
needs no LLM re-extraction.
Args:
store: The pgvector store, for its embeddings client.
conn: Open connection to the same database.
source_id: Source whose nodes to re-embed.
batch_size: Names per embed call and per transaction.
dry_run: When true, count the work and change nothing.
Returns:
Number of node rows seen (dry run) or rewritten.
"""
if not _graph_nodes_exist(conn):
return 0
cursor = conn.cursor()
try:
cursor.execute(
"SELECT id, name FROM graph_nodes "
"WHERE source_id = %s AND name_embedding IS NOT NULL ORDER BY id",
(source_id,),
)
rows = cursor.fetchall()
finally:
cursor.close()
if dry_run or not rows:
return len(rows)
written = 0
for batch in _batched(rows, batch_size):
vectors = store._embedding.embed_documents([row[1] or "" for row in batch])
cursor = conn.cursor()
try:
cursor.executemany(
"UPDATE graph_nodes SET name_embedding = %s::vector WHERE id = %s",
[
(str(list(vector)), row[0])
for vector, row in zip(vectors, batch)
],
)
conn.commit()
except Exception:
conn.rollback()
raise
finally:
cursor.close()
written += len(batch)
return written
def _count_source_chunks(conn, table: str, source_id: str) -> int:
"""How many chunk rows one source holds."""
cursor = conn.cursor()
try:
cursor.execute(
sql.SQL("SELECT count(*) FROM {table} WHERE source_id = %s").format(
table=sql.Identifier(table)
),
(source_id,),
)
return int(cursor.fetchone()[0])
finally:
cursor.close()
def _fetch_chunk_page(
conn, table: str, source_id: str, after_id: Optional[Any], limit: int
) -> List[Tuple[Any, str]]:
"""One page of ``(id, text)`` ordered by id, starting after ``after_id``.
Keyed off ``id`` rather than ``OFFSET`` so each page costs an index seek,
and the page boundaries stay stable across the updates the caller makes
between pages -- only the vector column changes, never ``id``.
"""
cursor = conn.cursor()
try:
if after_id is None:
cursor.execute(
sql.SQL(
"SELECT id, text FROM {table} WHERE source_id = %s "
"ORDER BY id LIMIT %s"
).format(table=sql.Identifier(table)),
(source_id, limit),
)
else:
cursor.execute(
sql.SQL(
"SELECT id, text FROM {table} WHERE source_id = %s AND id > %s "
"ORDER BY id LIMIT %s"
).format(table=sql.Identifier(table)),
(source_id, after_id, limit),
)
return cursor.fetchall()
finally:
cursor.close()
def _write_chunk_page(conn, table: str, vector_column: str, store, page) -> None:
"""Embed one page and commit its vectors."""
vectors = store._embedding.embed_documents([row[1] or "" for row in page])
cursor = conn.cursor()
try:
cursor.executemany(
sql.SQL("UPDATE {table} SET {column} = %s::vector WHERE id = %s").format(
table=sql.Identifier(table),
column=sql.Identifier(vector_column),
),
[(str(list(vector)), row[0]) for vector, row in zip(vectors, page)],
)
conn.commit()
except Exception:
conn.rollback()
raise
finally:
cursor.close()
def reembed_pgvector(source_id: str, batch_size: int, dry_run: bool) -> Tuple[int, int]:
"""Rewrite one source's vectors in place.
Rows are updated, never deleted and re-inserted, so an interrupted run
leaves every chunk present and simply re-does its batch next time.
Chunks are read a page at a time rather than all at once. A source's text is
the whole corpus: at the default chunk size 200k chunks materialise to about
1.6 GB of Python strings, several times that for non-Latin scripts, and the
result set is held alongside them until the cursor closes. That is enough to
be OOM-killed inside a 4 GiB container while also holding the model -- and a
kill here is what a torn index costs the most.
Args:
source_id: Source whose chunks to re-embed.
batch_size: Chunks per page, per embed call and per transaction.
dry_run: When true, count the work and change nothing.
Returns:
``(chunks_seen, chunks_written)``.
"""
store = VectorCreator.create_vectorstore("pgvector", source_id=source_id)
table, vector_column = store._table_name, store._vector_column
conn = store._get_connection()
seen = written = 0
try:
if dry_run:
seen = _count_source_chunks(conn, table, source_id)
else:
after_id: Optional[Any] = None
while True:
page = _fetch_chunk_page(
conn, table, source_id, after_id, batch_size
)
if not page:
break
_write_chunk_page(conn, table, vector_column, store, page)
after_id = page[-1][0]
seen += len(page)
written += len(page)
# The graph seeds every traversal from its own vectors, so leaving them
# in the old model's space is the same silent mismatch this script
# exists to remove -- and at equal widths nothing would report it.
if getattr(settings, "GRAPHRAG_ENABLED", False):
nodes = reembed_graph_nodes(store, conn, source_id, batch_size, dry_run)
if nodes:
logger.info(
" %s: %d graph node name(s)%s",
source_id,
nodes,
"" if dry_run else " re-embedded",
)
finally:
store.close()
return seen, written
def reembed_faiss(source_id: str, batch_size: int, dry_run: bool) -> Tuple[int, int]:
"""Rebuild one FAISS index from the chunk text in its sidecar.
A FAISS index is a flat array, so this rebuilds rather than updates. Every
new vector is computed *before* the old index is touched, so an interrupt
during embedding leaves the existing index intact -- and ``save_local``
writes each file to a temporary path and moves it into place, so an
interrupt during the write does too.
Args:
source_id: Source whose index to rebuild.
batch_size: Chunks per embed call.
dry_run: When true, count the work and change nothing.
Returns:
``(chunks_seen, chunks_written)``.
"""
# A width change is the main reason to run this, and it is exactly what
# ``assert_embedding_dimensions`` refuses to open. The chunk text lives in
# the sidecar rather than the index, so reading it needs no matching width,
# and the index this reads is replaced below.
store = VectorCreator.create_vectorstore(
"faiss",
source_id=source_id,
embeddings_key=settings.EMBEDDINGS_KEY,
skip_dimension_check=True,
)
chunks: List[Dict[str, Any]] = store.get_chunks() or []
if dry_run or not chunks:
return len(chunks), 0
docs = [
_RebuildDoc(chunk.get("text") or "", chunk.get("metadata") or {})
for chunk in chunks
]
# Constructing with ``docs_init`` embeds everything into a fresh in-memory
# index and touches nothing on disk. The existing index is only replaced by
# ``save_local`` below, so a failure while embedding leaves it intact.
# Keep the existing ids: GraphRAG's ``graph_node_chunks`` rows and any id a
# client already holds point at them.
ids = [chunk.get("doc_id") for chunk in chunks]
rebuilt = VectorCreator.create_vectorstore(
"faiss",
source_id=source_id,
embeddings_key=settings.EMBEDDINGS_KEY,
docs_init=docs,
ids=ids if all(ids) else None,
batch_size=batch_size,
)
rebuilt.save_local()
return len(chunks), len(docs)
def record_source_model(source_id: str) -> None:
"""Stamp ``sources.model`` with the model these vectors were just built by.
``sources`` lives in the user-data database while the vectors may not, so
this takes its own session. Left unwritten, the column keeps naming the old
model and every consumer that trusts it -- the boot mismatch check above
all -- reports a source as stale immediately after it was migrated.
"""
from sqlalchemy import text
from application.storage.db.session import db_session
try:
with db_session() as conn:
conn.execute(
text("UPDATE sources SET model = :model WHERE id = :id"),
{"model": settings.EMBEDDINGS_NAME, "id": source_id},
)
except Exception as exc: # noqa: BLE001 — the vectors are already rewritten
logger.warning(
" %s: re-embedded, but could not update sources.model (%s). Retrieval "
"is correct; the mismatch warning may persist until it is.",
source_id,
exc,
)
def run(
store_type: str,
source_ids: Optional[Sequence[str]],
batch_size: int,
dry_run: bool,
) -> int:
"""Re-embed the requested sources. Returns a process exit code."""
model = settings.EMBEDDINGS_NAME
spec = resolve(model)
logger.info(
"Re-embedding %s with %s (%s)",
store_type,
model,
f"{spec.dimension} dims" if spec else "width determined at runtime",
)
if dry_run:
logger.info("Dry run: nothing will be written.")
ids = list(source_ids) if source_ids else list_source_ids(store_type)
if not ids:
logger.warning("No sources found; nothing to do.")
return 0
logger.info("%d source(s) to process", len(ids))
handler = reembed_pgvector if store_type == "pgvector" else reembed_faiss
total_seen = total_written = 0
failures: List[str] = []
started = time.monotonic()
for position, source_id in enumerate(ids, start=1):
try:
seen, written = handler(source_id, batch_size, dry_run)
if written and not dry_run:
record_source_model(source_id)
total_seen += seen
total_written += written
logger.info(
"[%d/%d] %s: %d chunk(s)%s",
position,
len(ids),
source_id,
seen,
"" if dry_run else f", {written} re-embedded",
)
except Exception as exc: # noqa: BLE001 -- one bad source must not stop the run
failures.append(source_id)
logger.error("[%d/%d] %s FAILED: %s", position, len(ids), source_id, exc)
elapsed = time.monotonic() - started
logger.info(
"Done in %.1fs: %d chunk(s) seen, %d re-embedded, %d source(s) failed",
elapsed,
total_seen,
total_written,
len(failures),
)
if failures:
logger.error("Failed sources: %s", ", ".join(failures))
logger.error("Re-run with --sources %s to retry just those.", ",".join(failures))
return 1
if dry_run:
logger.info("Re-run without --dry-run to apply.")
return 0
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
prog="python -m application.scripts.reembed",
description=__doc__.split("\n\n")[0],
)
parser.add_argument(
"--dry-run", action="store_true", help="report what would change, write nothing"
)
parser.add_argument(
"--sources", help="comma-separated source ids; default is every source"
)
parser.add_argument(
"--batch-size",
type=int,
default=DEFAULT_BATCH_SIZE,
help=f"chunks per embed call (default {DEFAULT_BATCH_SIZE})",
)
parser.add_argument("-v", "--verbose", action="store_true", help="debug logging")
return parser
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
# 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(
"VECTOR_STORE is %r; this script supports %s. For other stores, "
"delete and re-ingest the affected sources.",
settings.VECTOR_STORE,
" and ".join(SUPPORTED_STORES),
)
return 2
sources = (
[s.strip() for s in args.sources.split(",") if s.strip()]
if args.sources
else None
)
try:
return run(store_type, sources, max(1, args.batch_size), args.dry_run)
except ReembedError as exc:
logger.error("%s", exc)
return 2
if __name__ == "__main__":
sys.exit(main())
+83 -20
View File
@@ -76,6 +76,43 @@ def ensure_database_ready(
_run_migrations(log)
def _release_boot_only_embeddings(log: logging.Logger) -> None:
"""Drop a model this hook loaded that the process will never use again.
``ensure_vector_schema`` has to run the model to learn the width of a model
the registry does not describe. ``EmbeddingsSingleton`` then caches it for
the life of the process -- correct when this process embeds, pure waste in
an API that delegates every embed to the worker, where it costs ~400 MB for
a small model and ~800 MB for a granite-sized one that is never called.
Only the local-ONNX case is dropped. A ``RemoteEmbeddings`` is cheap and is
exactly what the process goes on to use, and with delegation off the model
would only be rebuilt on the first query.
Bounds retention, not the transient peak: the load still happens, and the
ONNX Runtime arena may not return every page to the OS.
"""
from application.core.settings import settings
if settings.EMBEDDINGS_BASE_URL:
return
if getattr(settings, "EMBEDDINGS_DELEGATE_TO_WORKER", False) is not True:
return
import gc
from application.vectorstore.base import EmbeddingsSingleton
if EmbeddingsSingleton._instances.pop(settings.EMBEDDINGS_NAME, None) is None:
return
gc.collect()
log.info(
"ensure_vector_schema: released the embeddings model loaded to read its "
"width; this process delegates embedding to the worker and would never "
"have used it."
)
def ensure_vector_schema(*, logger: Optional[logging.Logger] = None) -> None:
"""Create the pgvector schema once at boot and verify its dimension.
@@ -127,27 +164,13 @@ 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,
)
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,
)
dim: Optional[int] = dimension_for(settings.EMBEDDINGS_NAME)
graph_enabled = bool(getattr(settings, "GRAPHRAG_ENABLED", False))
started = time.monotonic()
@@ -166,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 (
+150
View File
@@ -0,0 +1,150 @@
"""The embedding model an installation is pinned to.
``EMBEDDINGS_NAME`` has a code-level default, and moving that default would
re-point an existing index at a different vector space without anything
noticing: mpnet and granite are both 768-dimensional, so no width check fires
and retrieval simply gets worse. Which model an index was built with is a
property of the installation, not of the release it happens to be running.
So it is resolved once at boot and stored in ``app_metadata``. A fresh install
is pinned to the current recommendation; an install that already has sources is
pinned to the legacy model it has been using all along, and told how to move.
An explicit ``EMBEDDINGS_NAME`` in the environment always wins over both.
"""
from __future__ import annotations
import logging
from typing import Optional
from sqlalchemy import text
from application.core.settings import settings
from application.storage.db.repositories.app_metadata import AppMetadataRepository
from application.storage.db.session import db_session
from application.vectorstore.model_registry import (
DEFAULT_LEGACY,
DEFAULT_NEW_INSTALL,
resolve,
)
logger = logging.getLogger(__name__)
PIN_KEY = "embeddings_name"
NOTICE_KEY = "embeddings_legacy_notice_shown"
def _has_sources(conn) -> bool:
"""True when this installation already has an index to protect.
Counts ``sources`` rather than vector rows so the answer is the same for
every vector store, including the FAISS ones whose vectors are not in this
database at all.
"""
if conn.execute(text("SELECT to_regclass('public.sources')")).scalar() is None:
return False
return bool(conn.execute(text("SELECT EXISTS (SELECT 1 FROM sources)")).scalar())
def _legacy_notice(model: str) -> str:
return (
f"Embeddings: this installation is using {model}, the model its index was "
"built with, and will keep using it.\n"
"To move to granite (multilingual, a 32k-token context, same 768 dimensions):\n"
" 1. EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2\n"
" 2. python -m application.scripts.reembed\n"
"Changing the model without step 2 leaves queries searching a different "
"vector space than the stored vectors, which fails silently."
)
def resolve_embeddings_pin(log: Optional[logging.Logger] = None) -> None:
"""Point ``settings.EMBEDDINGS_NAME`` at this installation's pinned model.
Called once at boot, before anything embeds and before the vector schema
hook reads the width. Consumers read ``settings.EMBEDDINGS_NAME`` lazily, so
assigning it here is enough; no caller needs to know the pin exists.
Does nothing when the environment pins the name, and degrades to the
code-level default when the database cannot be reached.
Args:
log: Logger for the first-run notice; module logger when omitted.
"""
out = log or logger
if "EMBEDDINGS_NAME" in settings.model_fields_set:
return
try:
with db_session() as conn:
repo = AppMetadataRepository(conn)
stored = repo.get(PIN_KEY)
if stored:
settings.EMBEDDINGS_NAME = stored
return
upgrading = _has_sources(conn)
pinned = repo.setdefault(
PIN_KEY, DEFAULT_LEGACY if upgrading else DEFAULT_NEW_INSTALL
)
settings.EMBEDDINGS_NAME = pinned
out.info("Embeddings: pinned this installation to %s.", pinned)
if upgrading and repo.get(NOTICE_KEY) is None:
print(_legacy_notice(pinned), flush=True)
repo.set(NOTICE_KEY, "1")
except Exception as exc: # noqa: BLE001 — never block boot on the database
out.debug(
"Embeddings: could not resolve the pinned model (%s); using %s.",
exc,
settings.EMBEDDINGS_NAME,
exc_info=True,
)
def warn_on_source_model_mismatch(log: Optional[logging.Logger] = None) -> None:
"""Report sources whose vectors were built by a different model.
The width check cannot see this: mpnet and granite are both 768, so an
index queried by the wrong model returns worse answers and no error. This
is the only signal, so it names the sources and the command that fixes
them.
A source with no recorded model pre-dates the column and is therefore the
legacy model, not unknown.
Args:
log: Logger to warn on; module logger when omitted.
"""
out = log or logger
active = resolve(settings.EMBEDDINGS_NAME)
try:
with db_session() as conn:
if conn.execute(text("SELECT to_regclass('public.sources')")).scalar() is None:
return
rows = conn.execute(
text("SELECT COALESCE(model, :legacy) AS model, count(*) FROM sources GROUP BY 1"),
{"legacy": DEFAULT_LEGACY},
).fetchall()
except Exception as exc: # noqa: BLE001 — never block boot on the database
out.debug("Embeddings: could not check source models (%s)", exc, exc_info=True)
return
stale = [
(name, count)
for name, count in rows
# Compare through the registry so an alias is not read as a different
# model. An unregistered name resolves to None, which only matches
# another unregistered name if the strings agree.
for a, b in [(resolve(name), active)]
if (a or name) != (b or settings.EMBEDDINGS_NAME)
]
if not stale:
return
detail = ", ".join(f"{count} built with {name}" for name, count in stale)
out.warning(
"Embeddings: %s, but queries are embedded with %s. Retrieval against those "
"sources is degraded and will not raise. Re-embed them with "
"`python -m application.scripts.reembed`, or set EMBEDDINGS_NAME back.",
detail,
settings.EMBEDDINGS_NAME,
)
+6
View File
@@ -235,6 +235,12 @@ token_usage_table = Table(
# Added in ``0015_token_usage_model_id``. Canonical model id (catalog
# name for built-ins, UUID for BYOM); NULL on un-backfilled rows.
Column("model_id", Text),
# Added in ``0031_token_usage_cache_tokens``. Prompt-cache breakdowns of
# ``prompt_tokens`` as reported by the provider (reads; and writes on
# newer OpenAI-family models). NULL = not reported, distinct from 0 = no
# cache activity, so hit-rate queries stay honest across providers.
Column("cached_tokens", Integer),
Column("cache_write_tokens", Integer),
)
user_logs_table = Table(
@@ -37,6 +37,21 @@ class AppMetadataRepository:
{"key": key, "value": value},
)
def setdefault(self, key: str, value: str) -> str:
"""Store ``value`` under ``key`` only if absent, returning what stands.
Two workers racing on first startup converge on one value instead of
each overwriting the other, the way ``get_or_create_instance_id`` does.
"""
self._conn.execute(
text(
"INSERT INTO app_metadata (key, value) VALUES (:key, :value) "
"ON CONFLICT (key) DO NOTHING"
),
{"key": key, "value": value},
)
return self.get(key) or value
def get_or_create_instance_id(self) -> str:
"""Return the anonymous instance UUID, generating one if absent.
@@ -35,6 +35,8 @@ class TokenUsageRepository:
request_id: Optional[str] = None,
model_id: Optional[str] = None,
timestamp: Optional[datetime] = None,
cached_tokens: Optional[int] = None,
cache_write_tokens: Optional[int] = None,
) -> None:
# Attribution guard: the ``token_usage_attribution_chk`` CHECK
# constraint requires at least one of ``user_id`` / ``api_key``
@@ -60,12 +62,14 @@ class TokenUsageRepository:
INSERT INTO token_usage (
user_id, api_key, agent_id,
prompt_tokens, generated_tokens,
cached_tokens, cache_write_tokens,
source, request_id, model_id, timestamp
)
VALUES (
:user_id, :api_key,
CAST(:agent_id AS uuid),
:prompt_tokens, :generated_tokens,
:cached_tokens, :cache_write_tokens,
:source, :request_id, :model_id, COALESCE(:timestamp, now())
)
"""
@@ -76,6 +80,8 @@ class TokenUsageRepository:
"agent_id": agent_id_uuid,
"prompt_tokens": prompt_tokens,
"generated_tokens": generated_tokens,
"cached_tokens": cached_tokens,
"cache_write_tokens": cache_write_tokens,
"source": source,
"request_id": request_id,
"model_id": model_id,
+45 -8
View File
@@ -1,6 +1,7 @@
"""Local file system implementation."""
import os
import shutil
import tempfile
from typing import BinaryIO, List, Callable
from application.storage.base import BaseStorage
@@ -36,16 +37,52 @@ class LocalStorage(BaseStorage):
return resolved
def save_file(self, file_data: BinaryIO, path: str, **kwargs) -> dict:
"""Save a file to local storage."""
"""Save a file, replacing any existing one atomically.
The bytes land on a temporary file beside the destination and are moved
into place with ``os.replace``, so an interrupted write leaves the
previous file intact instead of a truncated one. Streaming straight onto
the destination is unrecoverable for a file that is rewritten in place:
a half-written ``index.faiss`` loads at neither the old width nor the
new one, and ``application.scripts.reembed`` rewrites every index it
touches.
"""
full_path = self._get_full_path(path)
directory = os.path.dirname(full_path)
os.makedirs(directory, exist_ok=True)
os.makedirs(os.path.dirname(full_path), exist_ok=True)
if hasattr(file_data, 'save'):
file_data.save(full_path)
else:
with open(full_path, 'wb') as f:
shutil.copyfileobj(file_data, f)
# Same directory as the destination, so the replace below is a rename
# within one filesystem rather than a copy.
fd, temp_path = tempfile.mkstemp(
dir=directory, prefix=f".{os.path.basename(full_path)}.", suffix=".tmp"
)
try:
with os.fdopen(fd, "wb") as f:
if hasattr(file_data, "save"):
file_data.save(f)
else:
shutil.copyfileobj(file_data, f)
f.flush()
os.fsync(f.fileno())
# mkstemp is 0600; keep whatever the file already had, or fall back
# to what a plain open() would have produced.
try:
mode = os.stat(full_path).st_mode & 0o777
except FileNotFoundError:
mode = 0o644
os.chmod(temp_path, mode)
os.replace(temp_path, full_path)
except BaseException:
# A successful replace consumes the temp file; every other exit
# leaves it behind.
try:
os.unlink(temp_path)
except OSError:
# Best-effort: the write already failed, and that exception is
# the one worth raising. A temp file we cannot remove must not
# mask it.
pass
raise
return {
'storage_type': 'local'
+169 -1
View File
@@ -2,12 +2,17 @@
from __future__ import annotations
import codecs
import io
import os
from contextlib import suppress
from typing import BinaryIO
from typing import BinaryIO, Container
from application.core.settings import settings
from application.parser.file.constants import (
attachment_extension,
has_attachment_parser,
)
_COPY_CHUNK_BYTES = 64 * 1024
@@ -26,6 +31,169 @@ class UploadTooLargeError(ValueError):
"""Raised when one uploaded file exceeds the configured byte cap."""
class UnsupportedUploadTypeError(ValueError):
"""Raised when an uploaded attachment has a file type the worker cannot parse."""
_UNSUPPORTED_UPLOAD_PREFIX = "Unsupported file type"
def unsupported_upload_message(filename: str | None) -> str:
"""Return the stable client-facing rejection message for an unparseable upload.
Args:
filename: The rejected upload's filename.
Returns:
``"Unsupported file type: .mp4"`` style text, naming the extension.
"""
extension = attachment_extension(filename)
return f"{_UNSUPPORTED_UPLOAD_PREFIX}: {extension or '(no extension)'}"
def is_unsupported_upload_message(message: str | None) -> bool:
"""Return whether ``message`` is one produced by :func:`unsupported_upload_message`."""
return bool(message) and str(message).startswith(_UNSUPPORTED_UPLOAD_PREFIX)
# A suffix with no parser is read by ``SimpleDirectoryReader``'s plain-text
# fallthrough, so it is admitted on content: enough of the head to recognise a
# container header, and a tolerance that keeps real text (UTF-8 accents, an
# ANSI-coloured log) in while a random binary — ~12.5% of bytes below 0x20 —
# stays out.
_TEXT_SNIFF_BYTES = 8192
_MAX_NONTEXT_RATIO = 0.10
# Control bytes that occur in ordinary text: tab, LF, VT, FF, CR, ESC.
_TEXT_CONTROL_BYTES = frozenset({0x09, 0x0A, 0x0B, 0x0C, 0x0D, 0x1B})
_TEXT_CONTROL_CHARS = frozenset("\t\n\v\f\r\x1b")
# A UTF-16/32 file is half NUL bytes, so the byte rules below would reject it.
# Its BOM says which encoding to read it as; the decoded characters are then
# judged instead. Longest BOM first — UTF-32-LE starts with the UTF-16-LE one.
# A BOM is never a verdict on its own: three bytes must not buy a video a pass.
_TEXT_BOMS = (
(codecs.BOM_UTF32_LE, "utf-32-le"),
(codecs.BOM_UTF32_BE, "utf-32-be"),
(codecs.BOM_UTF8, None),
(codecs.BOM_UTF16_LE, "utf-16-le"),
(codecs.BOM_UTF16_BE, "utf-16-be"),
)
def _decoded_looks_like_text(text: str) -> bool:
"""Return whether decoded characters read as text.
Args:
text: Characters decoded from a BOM-marked sample.
Returns:
False when a NUL character appears — the byte-level rule, one level
up, since no text holds one — or when too many characters are
unprintable: control, unassigned, private-use or surrogate, which is
what binary decodes into. Replacement characters count as unprintable,
so bytes the decoder could not read are evidence rather than silently
dropped (a truncated character at the sample boundary is one of
thousands and cannot swing the ratio).
"""
if not text:
return True
if "\x00" in text:
return False
nontext = sum(
1
for char in text
if char == "�"
or (not char.isprintable() and char not in _TEXT_CONTROL_CHARS)
)
return nontext / len(text) <= _MAX_NONTEXT_RATIO
def looks_like_text(sample: bytes) -> bool:
"""Return whether a leading byte sample reads as text rather than binary.
Args:
sample: The first bytes of a file; an empty sample counts as text.
Returns:
False when the sample holds a NUL byte or too many other non-text
control bytes, True otherwise. A UTF-16/32 BOM switches the test to
the decoded characters; the content is still what decides.
"""
if not sample:
return True
body = sample
for bom, encoding in _TEXT_BOMS:
if not sample.startswith(bom):
continue
body = sample[len(bom) :]
if encoding is not None:
return _decoded_looks_like_text(body.decode(encoding, errors="replace"))
# UTF-8 BOM: the byte rules still apply to everything after it.
break
if not body:
return True
if b"\x00" in body:
return False
nontext = sum(
1
for byte in body
if (byte < 0x20 and byte not in _TEXT_CONTROL_BYTES) or byte == 0x7F
)
return nontext / len(body) <= _MAX_NONTEXT_RATIO
def file_looks_like_text(path: str | os.PathLike[str]) -> bool:
"""Return whether a file on disk reads as text, by its leading bytes.
Args:
path: Filesystem path to sample.
Returns:
The :func:`looks_like_text` verdict for the file's head; True when the
file cannot be read, leaving that failure to the parser to report.
"""
try:
with open(path, "rb") as handle:
return looks_like_text(handle.read(_TEXT_SNIFF_BYTES))
except OSError:
return True
def enforce_parseable_attachment(
path: str | os.PathLike[str],
filename: str | None,
parser_extensions: Container[str] | None = None,
) -> None:
"""Reject an attachment that no parser handles and that is not plain text.
Suffixes with a parser are admitted unconditionally — a PDF is binary and
parses fine. Everything else, .txt included, has to read as text: that is
what keeps a video or an archive out of the plain-text fallthrough while
leaving source, config and log files in, whatever the file is named.
Args:
path: Local path to the staged upload, readable before it is stored.
filename: The upload's filename, used for the suffix and the message.
parser_extensions: Suffixes to treat as parser-backed. Defaults to
``ATTACHMENT_PARSER_EXTENSIONS``, which assumes the full parser
table; callers holding the live extractor (the worker) pass its
keys, so a docling-less install does not admit a .webp on trust
and then read it as text.
Raises:
UnsupportedUploadTypeError: When the file has no parser and its
contents are binary.
"""
if parser_extensions is None:
has_parser = has_attachment_parser(filename)
else:
has_parser = attachment_extension(filename) in parser_extensions
if has_parser:
return
if file_looks_like_text(path):
return
raise UnsupportedUploadTypeError(unsupported_upload_message(filename))
class _LimitedRawReader(io.RawIOBase):
"""Expose a binary stream while enforcing an exact cumulative read cap."""
+28 -1
View File
@@ -131,6 +131,11 @@ def _persist_call_usage(llm, call_usage):
agent_id=str(agent_id) if agent_id else None,
prompt_tokens=call_usage["prompt_tokens"],
generated_tokens=call_usage["generated_tokens"],
# Present only when the provider reported the breakdown;
# persisted as NULL otherwise so "unknown" never reads as
# "0% cache hits".
cached_tokens=call_usage.get("cached_tokens"),
cache_write_tokens=call_usage.get("cache_write_tokens"),
source=(
getattr(llm, "_token_usage_source", None) or "agent_stream"
),
@@ -151,6 +156,12 @@ def _prefer_provider_usage(llm: Any, call_usage: Dict[str, int]) -> Dict[str, in
``*_tokens_details`` breakdowns (``cached_tokens``,
``reasoning_tokens``) back out of these bins — that would break
parity with what providers bill.
The prompt-cache sub-bins ARE carried alongside (``cached_tokens``,
``cache_write_tokens``; Anthropic's ``cache_creation_tokens`` maps to
the latter) so persistence and the finish events can chart them. They
are added only when the provider reported them. The rest of
``call_usage`` (e.g. ``model``) is preserved rather than replaced.
"""
reported = getattr(llm, "_last_usage", None)
if not isinstance(reported, dict):
@@ -174,10 +185,22 @@ def _prefer_provider_usage(llm: Any, call_usage: Dict[str, int]) -> Dict[str, in
# Slotted/immutable LLM stand-ins can't record the claim; this
# call still gets the provider counts, which is correct for them.
pass
return {
merged = {
**call_usage,
"prompt_tokens": int(prompt or 0),
"generated_tokens": int(completion or 0),
}
details = reported.get("prompt_tokens_details")
if isinstance(details, dict):
cached = details.get("cached_tokens")
written = details.get("cache_write_tokens")
if written is None:
written = details.get("cache_creation_tokens")
if cached is not None:
merged["cached_tokens"] = int(cached or 0)
if written is not None:
merged["cache_write_tokens"] = int(written or 0)
return merged
def gen_token_usage(func):
@@ -224,6 +247,8 @@ def gen_token_usage(func):
prompt_tokens=call_usage["prompt_tokens"],
completion_tokens=call_usage["generated_tokens"],
latency_ms=int((time.monotonic() - started_at) * 1000),
cached_tokens=call_usage.get("cached_tokens"),
cache_write_tokens=call_usage.get("cache_write_tokens"),
error=error,
)
except Exception:
@@ -272,6 +297,8 @@ def stream_token_usage(func):
prompt_tokens=call_usage["prompt_tokens"],
completion_tokens=call_usage["generated_tokens"],
latency_ms=int((time.monotonic() - started_at) * 1000),
cached_tokens=call_usage.get("cached_tokens"),
cache_write_tokens=call_usage.get("cache_write_tokens"),
error=error,
)
except Exception:
+151 -77
View File
@@ -1,5 +1,4 @@
import logging
import os
from abc import ABC, abstractmethod
from typing import Optional
@@ -7,7 +6,29 @@ import requests
from application.core.settings import settings
from application.vectorstore.embeddings_openai import OpenAIEmbeddings
from application.utils import get_encoding
from application.vectorstore.model_registry import (
dimension_for,
max_input_tokens_for,
resolve,
)
def _embeddings_name_is_explicit() -> bool:
"""True when ``EMBEDDINGS_NAME`` names a model somebody actually chose.
Not ``model_fields_set``: pydantic marks a field as set for any value that
reached it, including one read from ``.env``, and every setup script has
always written ``EMBEDDINGS_NAME`` unconditionally. An install carrying the
legacy name a script wrote years ago would read as a deliberate choice and
lend a remote server that model's context window.
"""
fields = getattr(type(settings), "model_fields", None)
if not isinstance(fields, dict):
return False
field = fields.get("EMBEDDINGS_NAME")
if field is None:
return False
return settings.EMBEDDINGS_NAME != field.default
class RemoteEmbeddings:
@@ -23,45 +44,83 @@ class RemoteEmbeddings:
self.headers = {"Content-Type": "application/json"}
if api_key:
self.headers["Authorization"] = f"Bearer {api_key}"
self.dimension = 768
# Width comes from the registry. This used to be a hardcoded 768 that
# ``embed_query`` claimed to correct on first use -- but the correction
# was guarded by ``if self.dimension is None``, which the hardcode made
# unreachable, so a remote model of any other width silently produced a
# ``vector(768)`` column. ``None`` here means "unknown", and the probe
# below now genuinely runs.
self.dimension = dimension_for(model_name)
def _token_counter(self):
"""Counter matching the remote model's tokenizer, cached per process."""
from application.parser.tokenization import get_token_counter
return get_token_counter(self.model_name)
def _resolve_input_limit(self):
"""Token ceiling for a single embed input, or ``None`` for no limit.
``EMBEDDINGS_MAX_INPUT_TOKENS`` wins when set. Otherwise a registered
model contributes its own context window, so a request that the server
would reject -- or silently truncate -- is clipped here instead of
being sent and paid for.
That fallback needs the name to mean something. For a remote server it
is only a label forwarded as the ``model`` field, so the settings
default must not lend the server mpnet's 384-token window: a name
nobody chose describes nothing, and clipping on it would silently
discard most of every chunk.
"""
configured = settings.EMBEDDINGS_MAX_INPUT_TOKENS
if configured and configured > 0:
return configured
if not _embeddings_name_is_explicit():
return None
model_limit = max_input_tokens_for(self.model_name)
return model_limit if model_limit and model_limit > 0 else None
def _truncate_inputs(self, inputs):
"""Clip each input to ``EMBEDDINGS_MAX_INPUT_TOKENS`` tokens.
"""Clip each input to the resolved token limit.
The remote server (e.g. llama.cpp) hard-rejects any single input
larger than its physical batch size with a 500. When the setting is
configured, each input is truncated to that many tokens before the
request and the overflow is dropped (lossy by design). Token counts
use the shared tiktoken encoding, which differs from the server's
tokenizer, so set the limit with headroom under the server's true
limit to absorb tokenizer skew.
larger than its physical batch size with a 500, so oversized inputs are
truncated before the request and the overflow is dropped (lossy by
design).
Counting uses the embedding model's own tokenizer where it is known, so
the limit and the count are in the same unit. When it is not -- an
unregistered model, or no tokenizer available -- this falls back to
tiktoken, and the limit should then carry headroom to absorb the skew
between the two tokenizers.
Args:
inputs: A single string or a list of strings to embed.
Returns:
The inputs with each string clipped to the token limit, or the
inputs unchanged when the limit is unset or non-positive.
inputs unchanged when no limit applies.
"""
limit = settings.EMBEDDINGS_MAX_INPUT_TOKENS
if not limit or limit <= 0:
limit = self._resolve_input_limit()
if not limit:
return inputs
encoding = get_encoding()
counter = self._token_counter()
def clip(text):
if not isinstance(text, str):
return text
tokens = encoding.encode_ordinary(text)
if len(tokens) <= limit:
count = counter.count(text)
if count <= limit:
return text
logging.warning(
"Truncating remote embeddings input from %d to %d tokens (%d dropped)",
len(tokens),
count,
limit,
len(tokens) - limit,
count - limit,
)
return encoding.decode(tokens[:limit])
pieces = counter.split(text, limit)
return pieces[0] if pieces else text
if isinstance(inputs, list):
return [clip(text) for text in inputs]
@@ -129,7 +188,7 @@ class RemoteEmbeddings:
def _get_embeddings_wrapper():
"""Lazy import of EmbeddingsWrapper to avoid loading SentenceTransformer when using remote embeddings."""
"""Lazy import of EmbeddingsWrapper, so a remote setup never loads ONNX."""
from application.vectorstore.embeddings_local import EmbeddingsWrapper
return EmbeddingsWrapper
@@ -178,36 +237,28 @@ class EmbeddingsSingleton:
@staticmethod
def _create_instance(embeddings_name, *args, **kwargs):
if embeddings_name == "openai_text-embedding-ada-002":
"""Build the runner for ``embeddings_name``, per the model registry.
The registry replaced a hand-maintained factory dict whose entries
existed only to rewrite a configured name into a repository id. That
rewrite is now a registry field, so an unknown name needs no entry
here: it is passed through as a Hugging Face repository.
"""
spec = resolve(embeddings_name)
if spec is not None and spec.provider == "openai":
return OpenAIEmbeddings(*args, **kwargs)
# Lazy import EmbeddingsWrapper only when needed (avoids loading SentenceTransformer)
EmbeddingsWrapper = _get_embeddings_wrapper()
embeddings_factory = {
"huggingface_sentence-transformers/all-mpnet-base-v2": lambda: EmbeddingsWrapper(
"sentence-transformers/all-mpnet-base-v2"
),
"huggingface_sentence-transformers-all-mpnet-base-v2": lambda: EmbeddingsWrapper(
"sentence-transformers/all-mpnet-base-v2"
),
"huggingface_hkunlp/instructor-large": lambda: EmbeddingsWrapper(
"hkunlp/instructor-large"
),
}
if embeddings_name in embeddings_factory:
if args or kwargs:
logging.debug(
"Dropping %d positional and %d keyword argument(s) for pinned "
"embeddings model %s: its factory takes none.",
len(args),
len(kwargs),
embeddings_name,
)
return embeddings_factory[embeddings_name]()
else:
return EmbeddingsWrapper(embeddings_name, *args, **kwargs)
if spec is not None and (args or kwargs):
logging.debug(
"Dropping %d positional and %d keyword argument(s) for registered "
"embeddings model %s: the registry supplies its configuration.",
len(args),
len(kwargs),
embeddings_name,
)
return EmbeddingsWrapper(embeddings_name)
return EmbeddingsWrapper(embeddings_name, *args, **kwargs)
def _azure_configured() -> bool:
@@ -219,17 +270,29 @@ 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
):
"""Resolve the configured embeddings instance. The single entry point.
Callers that reach for :meth:`EmbeddingsSingleton.get_instance` directly
reproduce neither the bundled local-model path (the Docker image ships
``/app/models/all-mpnet-base-v2``) nor the OpenAI/Azure key handling: they
download a second copy of the model from the hub, and passing the key
positionally to the pinned HuggingFace names raises ``TypeError`` because
those factories take no arguments. Route every caller through here.
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``.
@@ -239,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
)
@@ -250,7 +341,11 @@ def get_embeddings(
)
return EmbeddingsSingleton._remote_instance(embeddings_name, embeddings_key)
if embeddings_name == "openai_text-embedding-ada-002":
# Match through the registry, not on the canonical string: the bare
# ``text-embedding-ada-002`` alias resolves here too, and skipping this
# branch would drop the Azure deployment name and the API key.
spec = resolve(embeddings_name)
if spec is not None and spec.provider == "openai":
if _azure_configured():
embedding_instance = EmbeddingsSingleton.get_instance(
embeddings_name, model=settings.AZURE_EMBEDDINGS_DEPLOYMENT_NAME
@@ -259,31 +354,10 @@ def get_embeddings(
embedding_instance = EmbeddingsSingleton.get_instance(
embeddings_name, openai_api_key=embeddings_key
)
elif embeddings_name == "huggingface_sentence-transformers/all-mpnet-base-v2":
possible_paths = [
"/app/models/all-mpnet-base-v2", # Docker absolute path
"./models/all-mpnet-base-v2", # Relative path
]
local_model_path = None
for path in possible_paths:
if os.path.exists(path):
local_model_path = path
logging.info(f"Found local model at path: {path}")
break
else:
logging.info(f"Path does not exist: {path}")
if local_model_path:
embedding_instance = EmbeddingsSingleton.get_instance(
local_model_path,
)
else:
logging.warning(
f"Local model not found in any of the paths: {possible_paths}. Falling back to HuggingFace download."
)
embedding_instance = EmbeddingsSingleton.get_instance(
embeddings_name,
)
else:
# No per-model branching: the registry resolves names and FastEmbed
# caches artifacts under EMBEDDINGS_CACHE_DIR, which is where the
# image warms them at build time.
embedding_instance = EmbeddingsSingleton.get_instance(embeddings_name)
return embedding_instance
@@ -0,0 +1,254 @@
"""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
import threading
import time
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"
#: 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
#: How long a caller waits for the outcome of the dispatch already in flight
#: before giving up on its own. Only applies while the worker is unproven --
#: once one dispatch has succeeded, every caller goes straight to the broker.
#: Comfortably above a healthy round trip (~60 ms on a prefork worker) and far
#: below ``EMBEDDINGS_DELEGATE_TIMEOUT``, which is the point.
_PROBE_WAIT = 2.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."""
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)
self._failed_at: Optional[float] = None
# A dispatch has completed successfully, so the worker is known to be
# consuming the queue and callers need not take turns proving it.
self._verified = False
self._probing = False
self._state_lock = threading.Lock()
self._probe_done = threading.Event()
def _cooldown_remaining(self) -> float:
"""Seconds left of the fail-fast window after a failed dispatch."""
# One load: a concurrent success clearing the latch between two reads
# would otherwise subtract from None.
failed_at = self._failed_at
if failed_at is None:
return 0.0
return max(0.0, _FAILURE_COOLDOWN - (time.monotonic() - failed_at))
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 _send(self, texts: List[str], queue: str, timeout: int) -> List[List[float]]:
"""Publish the embed task and wait for its vectors."""
from application.celery_init import celery
result = celery.send_task(EMBED_TASK, args=[texts, self.embeddings_name], queue=queue)
try:
vectors = result.get(timeout=timeout)
except Exception as exc:
self._failed_at = time.monotonic()
# Drop the proof with the worker that supplied it. ``_verified``
# short-circuits ahead of the probe gate, so leaving it set means
# the gate only ever covers a worker that was never healthy --
# while the case that actually happens is a healthy one being
# redeployed or OOM-killed. Every caller would then pay the full
# timeout, together, on every wave once the cooldown lapses.
self._verified = False
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 "
f"work. {_NO_WORKER_HINT}"
) from exc
finally:
_forget(result)
self._failed_at = None
self._verified = True
return vectors
def _cooldown_error(self, queue: str, remaining: float) -> RuntimeError:
return 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}"
)
def _dispatch(self, texts: List[str]) -> List[List[float]]:
"""Run the embed task on the worker and wait for its vectors.
A missing worker is a property of the deployment, not of this call, so
at most one caller waits out ``EMBEDDINGS_DELEGATE_TIMEOUT`` to discover
it. Two guards do that:
The cooldown latch covers requests arriving *after* a failure -- without
it a single retrieval pays the timeout twice, once in
``fanout.embed_questions`` and again per source when it falls back to
letting each store embed its own query.
The probe covers requests already in flight *alongside* the first one,
which the latch cannot: nothing is latched until that first ``get()``
returns, so every thread in the opening wave would otherwise block for
the full timeout at once -- at the shipped 60s and 96 WSGI threads, an
API that serves nothing at all, health checks included.
"""
queue = getattr(settings, "EMBEDDINGS_QUEUE", "embeddings")
timeout = getattr(settings, "EMBEDDINGS_DELEGATE_TIMEOUT", 60)
remaining = self._cooldown_remaining()
if remaining > 0:
raise self._cooldown_error(queue, remaining)
if self._verified:
return self._send(texts, queue, timeout)
with self._state_lock:
# "send" -- proven while we waited for the lock, just go.
# "wait" -- another caller is already finding out; don't pay a
# second full timeout to learn the same thing.
# "probe" -- nobody is; this call is the one that finds out.
role = "send" if self._verified else "wait" if self._probing else "probe"
if role == "probe":
self._probing = True
self._probe_done.clear()
if role == "probe":
try:
return self._send(texts, queue, timeout)
finally:
with self._state_lock:
self._probing = False
self._probe_done.set()
if role == "wait":
self._probe_done.wait(_PROBE_WAIT)
remaining = self._cooldown_remaining()
if remaining > 0:
raise self._cooldown_error(queue, remaining)
if not self._verified:
raise RuntimeError(
f"Skipping the embed dispatch: an earlier request to the "
f"{queue!r} queue is still unanswered after {_PROBE_WAIT}s, so "
f"no worker appears to be consuming it. {_NO_WORKER_HINT}"
)
return self._send(texts, queue, timeout)
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")
+336 -38
View File
@@ -1,54 +1,352 @@
"""
Local embeddings using SentenceTransformer.
This module is only imported when EMBEDDINGS_BASE_URL is not set,
to avoid loading SentenceTransformer into memory when using remote embeddings.
"""Local embeddings via FastEmbed (ONNX Runtime).
Replaces the previous SentenceTransformer implementation. Both run the same
weights; FastEmbed reaches them through ONNX Runtime instead of torch, which
removes torch, transformers and sentence-transformers from the dependency set
and measurably reduces both resident memory and import cost.
The swap is numerically transparent for existing indexes: embedding the same
text with SentenceTransformer and with FastEmbed's fp32 ONNX graph of mpnet
returns vectors at cosine 1.0, so an mpnet index built before this change
keeps working unchanged. That result is specific to the fp32 graph -- the
granite entries run an int8-quantised one and are not bit-comparable to a
fp32 index of the same model.
Models are described in :mod:`application.vectorstore.model_registry`. A name
the registry does not know is treated as a Hugging Face repository, which is
what someone configuring an arbitrary model expects; how to run it is read
from the repository itself rather than assumed.
"""
import json
import logging
import threading
from dataclasses import replace
from typing import Any, List, Optional
from sentence_transformers import SentenceTransformer
from application.core.settings import settings
from application.vectorstore.model_registry import EmbeddingModel, resolve, known_names
logger = logging.getLogger(__name__)
# ``add_custom_model`` mutates a process-global registry inside FastEmbed, so
# repeated registration of the same name is both wasteful and racy under the
# thread pool the API serves requests from.
_registered: set = set()
_register_lock = threading.Lock()
# Last-resort layout for a repository that declares nothing about itself.
_FALLBACK_ONNX_FILE = "onnx/model.onnx"
_FALLBACK_POOLING = "mean"
# Sentence-transformers records how a model turns token vectors into one
# vector, and whether it normalises the result, as files in the repository.
# Reading them is the difference between running a model and running something
# that merely shares its weights: mean-pooling a CLS model returns vectors at
# cosine ~0.95 to the correct ones -- close enough to look like it works, far
# enough to degrade retrieval, and silent either way.
_POOLING_CONFIG = "1_Pooling/config.json"
_MODULES_CONFIG = "modules.json"
def _pooling_type(pooling: str):
"""Map our ``"cls"``/``"mean"`` spelling onto FastEmbed's enum."""
from fastembed.common.model_description import PoolingType
return {"cls": PoolingType.CLS, "mean": PoolingType.MEAN}[pooling]
def _is_builtin(repo: str) -> bool:
"""True when FastEmbed already ships a description for ``repo``."""
from fastembed import TextEmbedding
lowered = repo.lower()
return any(
str(entry.get("model", "")).lower() == lowered
for entry in TextEmbedding.list_supported_models()
)
def _register(model: EmbeddingModel) -> None:
"""Teach FastEmbed about a model, exactly once per process."""
from fastembed import TextEmbedding
from fastembed.common.model_description import ModelSource
with _register_lock:
if model.repo in _registered:
return
if _is_builtin(model.repo):
# ``add_custom_model`` refuses a name FastEmbed already ships, and
# its own description carries the pooling, width and graph file we
# would be supplying, so there is nothing to add. Without this,
# configuring any of FastEmbed's ~30 built-in models (bge, e5,
# MiniLM, gte, ...) fails every embed call.
_registered.add(model.repo)
return
TextEmbedding.add_custom_model(
model=model.repo,
pooling=_pooling_type(model.pooling),
normalization=model.normalize,
sources=ModelSource(hf=model.repo),
dim=model.dimension,
model_file=model.onnx_file,
)
_registered.add(model.repo)
def _read_repo_json(repo: str, filename: str) -> Optional[dict]:
"""Fetch one small JSON from ``repo``, or ``None`` when it is not there.
Reads through the Hugging Face hub cache, so a warmed image finds it
offline. Every failure -- absent file, no network, malformed JSON -- is the
same answer to the caller: this repository does not tell us.
"""
try:
from huggingface_hub import hf_hub_download
with open(hf_hub_download(repo_id=repo, filename=filename), encoding="utf-8") as handle:
return json.load(handle)
except Exception as exc:
logger.debug("No %s for %s (%s)", filename, repo, exc)
return None
def _describe_from_repo(repo: str) -> Optional[EmbeddingModel]:
"""Build a spec from a repository's sentence-transformers metadata.
Args:
repo: Hugging Face repository id.
Returns:
The described model, or ``None`` when the repository carries no
metadata to read -- leaving the caller to fall back to assumptions.
Raises:
RuntimeError: If the model has a Dense projection head. FastEmbed runs
the transformer and pools it, and nothing else, so the projection
would be skipped and the vectors come out both the wrong width and
in a different space. There is no correct way to run it here.
"""
pooling_config = _read_repo_json(repo, _POOLING_CONFIG)
if pooling_config is None:
return None
kinds = {
str(module.get("type", "")).rsplit(".", 1)[-1]
for module in (_read_repo_json(repo, _MODULES_CONFIG) or [])
if isinstance(module, dict)
}
if "Dense" in kinds:
raise RuntimeError(
f"Embedding model {repo!r} has a Dense projection layer, which FastEmbed "
"cannot run: its vectors would be the wrong width and in a different "
"space than the model was trained to produce. Choose a model without "
"one, or serve this one over EMBEDDINGS_BASE_URL."
)
if pooling_config.get("pooling_mode_cls_token"):
pooling = "cls"
elif pooling_config.get("pooling_mode_mean_tokens"):
pooling = "mean"
else:
# max, mean_sqrt_len, weighted-mean: FastEmbed offers none of them, so
# there is nothing to describe and guessing is what we are avoiding.
logger.warning(
"Embedding model %s uses a pooling mode FastEmbed cannot reproduce (%s).",
repo,
", ".join(sorted(k for k, v in pooling_config.items() if v is True)) or "unknown",
)
return None
# ``Normalize`` appears in modules.json only when the model L2-normalises.
# Its absence is a fact, not missing data: a dot-product model is trained
# on unnormalised vectors and normalising re-ranks its results.
normalize = "Normalize" in kinds
dimension = int(pooling_config.get("word_embedding_dimension") or 0)
logger.info(
"Embedding model %s declares %s pooling, normalize=%s, dimension=%s.",
repo,
pooling,
normalize,
dimension or "unknown",
)
return EmbeddingModel(
name=repo,
dimension=dimension,
max_input_tokens=512,
pooling=pooling,
normalize=normalize,
repo=repo,
onnx_file=_FALLBACK_ONNX_FILE,
)
def _apply_overrides(spec: EmbeddingModel) -> EmbeddingModel:
"""Let ``EMBEDDINGS_POOLING``/``EMBEDDINGS_NORMALIZE`` win over any source."""
pooling = getattr(settings, "EMBEDDINGS_POOLING", None)
normalize = getattr(settings, "EMBEDDINGS_NORMALIZE", None)
changes = {}
if isinstance(pooling, str) and pooling.strip().lower() in ("cls", "mean"):
changes["pooling"] = pooling.strip().lower()
if isinstance(normalize, bool):
changes["normalize"] = normalize
if not changes:
return spec
logger.info("Overriding %s from settings: %s", spec.repo, changes)
return replace(spec, **changes)
def _spec_for(model_name: str) -> EmbeddingModel:
"""Registry entry for ``model_name``, or a best-effort one for a raw repo.
Args:
model_name: Configured ``EMBEDDINGS_NAME``.
Returns:
The registry entry, else one read from the repository's own
sentence-transformers metadata, else a last-resort entry assuming the
standard ONNX layout with mean pooling and ``dimension = 0`` so the
caller knows to probe for the real width. Settings overrides win over
all three.
"""
spec = resolve(model_name)
if spec is not None:
return _apply_overrides(spec)
described = _describe_from_repo(model_name)
if described is not None:
return _apply_overrides(described)
logger.warning(
"Embedding model %r is not in the registry (known: %s) and its repository "
"declares no pooling, so %s with L2 normalisation is assumed. If that is "
"wrong the vectors will be quietly poor rather than fail; set "
"EMBEDDINGS_POOLING=cls|mean and EMBEDDINGS_NORMALIZE to pin it.",
model_name,
", ".join(known_names()),
_FALLBACK_POOLING,
)
return _apply_overrides(
EmbeddingModel(
name=model_name,
dimension=0,
max_input_tokens=512,
pooling=_FALLBACK_POOLING,
normalize=True,
repo=model_name,
onnx_file=_FALLBACK_ONNX_FILE,
)
)
def _pad_to_longest_in_batch(model: Any) -> None:
"""Undo a fixed padding width baked into a model's ``tokenizer.json``."""
# FastEmbed enables padding only when the tokenizer declares none, so a
# fixed ``length`` survives loading. Shorter inputs are then padded to that
# width while longer ones keep their own, the batch is ragged, and the ONNX
# tensor build fails. mpnet ships ``length: 128``; granite does not.
# Mean pooling masks pad tokens, so the vectors are unaffected.
tokenizer = getattr(getattr(model, "model", None), "tokenizer", None)
padding = getattr(tokenizer, "padding", None)
if not isinstance(padding, dict) or padding.get("length") is None:
return
tokenizer.enable_padding(
direction=padding.get("direction", "right"),
pad_id=padding.get("pad_id", 0),
pad_type_id=padding.get("pad_type_id", 0),
pad_token=padding.get("pad_token", "<pad>"),
length=None,
pad_to_multiple_of=padding.get("pad_to_multiple_of"),
)
class EmbeddingsWrapper:
def __init__(self, model_name, *args, **kwargs):
logging.info(f"Initializing EmbeddingsWrapper with model: {model_name}")
"""Runs an embedding model locally through FastEmbed.
Exposes the ``embed_query``/``embed_documents``/``dimension`` interface the
vector stores rely on, matching ``RemoteEmbeddings`` and ``OpenAIEmbeddings``.
"""
def __init__(self, model_name: str, *args: Any, **kwargs: Any) -> None:
"""Load ``model_name`` locally.
Args:
model_name: Registry name, alias, or a Hugging Face repository id.
Raises:
RuntimeError: If the model cannot be loaded, with the configured
name and the registry's known names in the message.
"""
from fastembed import TextEmbedding
self.spec = _spec_for(model_name)
logger.info("Loading embeddings model %s via FastEmbed", self.spec.repo)
try:
kwargs.setdefault("trust_remote_code", True)
self.model = SentenceTransformer(
model_name,
config_kwargs={"allow_dangerous_deserialization": True},
*args,
**kwargs,
)
if self.model is None or self.model._first_module() is None:
raise ValueError(
f"SentenceTransformer model failed to load properly for: {model_name}"
)
# Renamed in sentence-transformers 5.4; keep the old name as a fallback.
get_dimension = getattr(
self.model,
"get_embedding_dimension",
getattr(self.model, "get_sentence_embedding_dimension", None),
)
self.dimension = get_dimension()
logging.info(f"Successfully loaded model with dimension: {self.dimension}")
except Exception as e:
logging.error(
f"Failed to initialize SentenceTransformer with model {model_name}: {str(e)}",
exc_info=True,
)
raise
_register(self.spec)
init_kwargs = {"model_name": self.spec.repo}
threads = getattr(settings, "EMBEDDINGS_THREADS", None)
if isinstance(threads, int) and threads > 0:
init_kwargs["threads"] = threads
cache_dir = getattr(settings, "EMBEDDINGS_CACHE_DIR", None)
if cache_dir:
init_kwargs["cache_dir"] = cache_dir
self.model = TextEmbedding(**init_kwargs)
except Exception as exc:
raise RuntimeError(
f"Could not load embeddings model {model_name!r} via FastEmbed: "
f"{exc}. Known models: {', '.join(known_names())}."
) from exc
def embed_query(self, query: str):
return self.model.encode(query).tolist()
_pad_to_longest_in_batch(self.model)
self.dimension = self.spec.dimension or self._probe_dimension()
logger.info("Embeddings model ready (dimension=%d)", self.dimension)
def embed_documents(self, documents: list):
return self.model.encode(documents).tolist()
def _probe_dimension(self) -> int:
"""Determine the vector width of a model the registry does not describe."""
return len(self.embed_query("dimension probe"))
def embed_query(self, query: str) -> List[float]:
"""Embed a single query string."""
return self.embed_documents([query])[0]
def embed_documents(self, documents: List[str]) -> List[List[float]]:
"""Embed a list of documents, preserving input order.
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.
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_MODEL_BATCH_SIZE", None)
if isinstance(raw, int) and not isinstance(raw, bool) and raw > 0:
batch_size = raw
documents = list(documents)
if batch_size is None or len(documents) <= batch_size:
# One batch either way: sorting would only add work.
return [v.tolist() for v in self.model.embed(documents, batch_size=batch_size)]
order = sorted(range(len(documents)), key=lambda i: len(documents[i]))
grouped = [documents[i] for i in order]
vectors = [v.tolist() for v in self.model.embed(grouped, batch_size=batch_size)]
restored: List[Optional[List[float]]] = [None] * len(documents)
for position, original_index in enumerate(order):
restored[original_index] = vectors[position]
return restored
def __call__(self, text):
if isinstance(text, str):
return self.embed_query(text)
elif isinstance(text, list):
return self.embed_documents(text)
else:
raise ValueError("Input must be a string or a list of strings")
raise ValueError("Input must be a string or a list of strings")
@@ -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))
+80 -27
View File
@@ -81,7 +81,29 @@ class FaissStore(BaseVectorStore):
# and must not be shown as one.
score_kind = "l2_distance"
def __init__(self, source_id: str, embeddings_key: str, docs_init=None):
def __init__(
self,
source_id: str,
embeddings_key: str,
docs_init=None,
ids=None,
batch_size=None,
skip_dimension_check: bool = False,
):
"""Open or build one source's FAISS index.
Args:
source_id: Source whose index to open.
embeddings_key: API key handed to the embeddings provider.
docs_init: Documents to build a fresh index from. Loads the stored
index instead when omitted.
ids: Chunk ids to keep when building. Generated when omitted.
batch_size: Documents per embed call when building.
skip_dimension_check: Open an index whose width does not match the
configured model. Only for a caller that is about to replace
that index, such as the re-embed script, which otherwise cannot
read the chunks it needs to rebuild from.
"""
super().__init__()
self.source_id = source_id
self.path = get_vectorstore(source_id)
@@ -94,27 +116,47 @@ class FaissStore(BaseVectorStore):
try:
if docs_init:
self._build_from_documents(docs_init)
self._build_from_documents(docs_init, ids=ids, batch_size=batch_size)
else:
self._load_from_storage()
except Exception as e:
raise Exception(f"Error loading FAISS index: {str(e)}")
self.assert_embedding_dimensions(self.embeddings)
if not skip_dimension_check:
self.assert_embedding_dimensions(self.embeddings)
# -- Construction ----------------------------------------------------
def _build_from_documents(self, docs_init) -> None:
"""Create a fresh index seeded with ``docs_init``."""
def _build_from_documents(self, docs_init, ids=None, batch_size=None) -> None:
"""Create a fresh index seeded with ``docs_init``.
Args:
docs_init: Documents to embed.
ids: Chunk ids to keep. Generated when omitted, which renumbers
every chunk and orphans anything referencing the old ids.
batch_size: Documents per embed call. Without it the whole index
goes out in one call, which a remote embeddings server rejects
or times out on.
"""
texts, metadatas = [], []
for doc in docs_init:
texts.append(getattr(doc, "page_content", None) or getattr(doc, "text", "") or "")
metadatas.append(getattr(doc, "metadata", None) or getattr(doc, "extra_info", None) or {})
faiss = _dependable_faiss_import()
vectors = self.embeddings.embed_documents(texts)
self.index = faiss.IndexFlatL2(len(vectors[0]))
self._append(texts, metadatas, vectors)
ids = list(ids) if ids else None
step = batch_size if batch_size and batch_size > 0 else len(texts)
for start in range(0, len(texts), max(1, step)):
stop = start + max(1, step)
vectors = self.embeddings.embed_documents(texts[start:stop])
if self.index is None:
self.index = faiss.IndexFlatL2(len(vectors[0]))
self._append(
texts[start:stop],
metadatas[start:stop],
vectors,
ids[start:stop] if ids else None,
)
def _load_from_storage(self) -> None:
"""Load the index and its sidecar, preferring JSON over the pickle."""
@@ -283,7 +325,13 @@ class FaissStore(BaseVectorStore):
f.write(dump_pickle_sidecar(self.documents, self.index_to_docstore_id))
def _save_to_storage(self) -> bool:
"""Persist the index through the configured storage backend."""
"""Persist the index through the configured storage backend.
Each file is replaced atomically by the backend, so an interrupted save
leaves the previous one readable. The three are still written in
sequence: a crash between them pairs a new index with an older sidecar,
which stays consistent because the ids and row order are preserved.
"""
with tempfile.TemporaryDirectory() as temp_dir:
self._write_index_files(temp_dir)
storage_path = get_vectorstore(self.source_id)
@@ -301,24 +349,29 @@ class FaissStore(BaseVectorStore):
# -- Introspection ---------------------------------------------------
def assert_embedding_dimensions(self, embeddings) -> None:
"""Check the index width matches the embedding model's width."""
if (
settings.EMBEDDINGS_NAME
== "huggingface_sentence-transformers/all-mpnet-base-v2"
):
word_embedding_dimension = getattr(embeddings, "dimension", None)
if word_embedding_dimension is None:
raise AttributeError(
"'dimension' attribute not found in embeddings instance."
)
if self.index is None:
return
if word_embedding_dimension != self.index.d:
raise ValueError(
f"Embedding dimension mismatch: embeddings.dimension "
f"({word_embedding_dimension}) != docsearch index dimension "
f"({self.index.d})"
)
"""Check the index width matches the embedding model's width.
This used to run only when ``EMBEDDINGS_NAME`` was mpnet, so every
other model skipped the check entirely -- exactly the models most
likely to differ from an index built earlier. It now runs for any
model that reports a width.
"""
word_embedding_dimension = getattr(embeddings, "dimension", None)
if word_embedding_dimension is None:
# A remote model of unknown width reports None until its first
# call; there is nothing to compare yet.
return
if self.index is None:
return
if word_embedding_dimension != self.index.d:
raise ValueError(
f"Embedding dimension mismatch: {settings.EMBEDDINGS_NAME} produces "
f"{word_embedding_dimension}-dim vectors but this FAISS index is "
f"{self.index.d}-dim. The index was built with a different "
f"embedding model; re-embed it with "
f"`python -m application.scripts.reembed` or point "
f"EMBEDDINGS_NAME back at the original model."
)
def get_chunks(self) -> List[Dict[str, Any]]:
"""Return every chunk held in the index."""
+5 -1
View File
@@ -2,6 +2,7 @@ from typing import List, Optional
import importlib
from application.vectorstore.base import BaseVectorStore
from application.core.settings import settings
from application.vectorstore.model_registry import DEFAULT_EMBEDDING_DIMENSION
class LanceDBVectorStore(BaseVectorStore):
"""Class for LanceDB Vector Store integration."""
@@ -54,8 +55,11 @@ class LanceDBVectorStore(BaseVectorStore):
"""Ensure the table exists before performing operations."""
if self.table is None:
embeddings = self._get_embeddings(settings.EMBEDDINGS_NAME, self.embeddings_key)
# A model outside the registry reports no width until it has run;
# ``list_size=None`` is a TypeError, not a permissive schema.
dimension = getattr(embeddings, "dimension", None) or DEFAULT_EMBEDDING_DIMENSION
schema = self.pa.schema([
self.pa.field("vector", self.pa.list_(self.pa.float32(), list_size=embeddings.dimension)),
self.pa.field("vector", self.pa.list_(self.pa.float32(), list_size=dimension)),
self.pa.field("text", self.pa.string()),
self.pa.field("metadata", self.pa.struct([
self.pa.field("key", self.pa.string()),
+177
View File
@@ -0,0 +1,177 @@
"""Canonical description of every embedding model DocsGPT knows how to run.
``EMBEDDINGS_NAME`` used to be a free-form string interpreted in half a dozen
places: a factory dict here, a bundled-model path probe there, a dimension
assertion gated on one model's name, and a hardcoded ``dimension = 768`` on the
remote client. Each of those encoded a different subset of the same facts, and
they drifted.
This module is the single place those facts live. A model is described once and
every consumer -- the local runner, the remote client, the schema bootstrap, the
chunker -- reads the same entry.
Unknown names are not an error: :func:`resolve` returns ``None`` and callers
fall back to treating the name as a Hugging Face repository, which is what a
user configuring an arbitrary model expects.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Dict, Optional, Tuple
@dataclass(frozen=True)
class EmbeddingModel:
"""Everything the application needs to know about one embedding model.
Attributes:
name: Canonical ``EMBEDDINGS_NAME`` value.
dimension: Width of the vectors it produces.
max_input_tokens: The model's own context window, counted in *its*
tokenizer. Used to bound what we send it, never to silently
reshape chunks.
pooling: ``"cls"`` or ``"mean"`` -- how token vectors become one vector.
normalize: Whether outputs are L2-normalised to unit length.
provider: Which runner handles it (``"fastembed"`` or ``"openai"``).
repo: Hugging Face repository holding weights and tokenizer.
onnx_file: Path within ``repo`` to the ONNX graph to run.
aliases: Other spellings of ``name`` accepted from configuration.
"""
name: str
dimension: int
max_input_tokens: int
pooling: str = "mean"
normalize: bool = True
provider: str = "fastembed"
repo: Optional[str] = None
onnx_file: Optional[str] = None
aliases: Tuple[str, ...] = field(default_factory=tuple)
#: The model DocsGPT installed before the granite migration. Kept as the
#: default so an existing deployment that upgrades keeps its index working;
#: new installs are pointed at granite by the setup script and env template.
MPNET = EmbeddingModel(
name="huggingface_sentence-transformers/all-mpnet-base-v2",
dimension=768,
max_input_tokens=384,
pooling="mean",
normalize=True,
repo="sentence-transformers/all-mpnet-base-v2",
onnx_file="onnx/model.onnx",
aliases=(
"huggingface_sentence-transformers-all-mpnet-base-v2",
"sentence-transformers/all-mpnet-base-v2",
"all-mpnet-base-v2",
),
)
#: Default for new installs: same 768 dimensions as mpnet, an 8x wider
#: effective input, and multilingual retrieval.
#:
#: Runs the int8-quantised graph, which is ~4x smaller than the fp32 one
#: (313 MB against 1247 MB) and keeps the image shippable now that both
#: defaults are baked in. It costs some numeric fidelity against fp32 --
#: measured at cosine 0.96 on short texts, with retrieval rank order
#: unchanged -- so it is not bit-comparable to a fp32 granite index.
GRANITE_311M = EmbeddingModel(
name="ibm-granite/granite-embedding-311m-multilingual-r2",
dimension=768,
max_input_tokens=32768,
pooling="cls",
normalize=True,
repo="ibm-granite/granite-embedding-311m-multilingual-r2",
onnx_file="onnx/model_quint8_avx2.onnx",
aliases=("granite-embedding-311m-multilingual-r2", "granite-311m"),
)
#: Smaller granite. Half the vector width, roughly three times the speed.
#: Int8-quantised on the same terms as GRANITE_311M above.
GRANITE_97M = EmbeddingModel(
name="ibm-granite/granite-embedding-97m-multilingual-r2",
dimension=384,
max_input_tokens=32768,
pooling="cls",
normalize=True,
repo="ibm-granite/granite-embedding-97m-multilingual-r2",
onnx_file="onnx/model_quint8_avx2.onnx",
aliases=("granite-embedding-97m-multilingual-r2", "granite-97m"),
)
OPENAI_ADA_002 = EmbeddingModel(
name="openai_text-embedding-ada-002",
dimension=1536,
max_input_tokens=8191,
pooling="mean",
normalize=True,
provider="openai",
repo=None,
onnx_file=None,
aliases=("text-embedding-ada-002",),
)
MODELS: Tuple[EmbeddingModel, ...] = (
MPNET,
GRANITE_311M,
GRANITE_97M,
OPENAI_ADA_002,
)
#: Fallback vector width when the configured model is unknown -- the width
#: every DocsGPT install has used to date, so an unrecognised model does not
#: silently reshape an existing table.
DEFAULT_EMBEDDING_DIMENSION = 768
#: Name that a fresh install should be configured with.
DEFAULT_NEW_INSTALL = GRANITE_311M.name
#: Name that ``settings.EMBEDDINGS_NAME`` defaults to, i.e. what an existing
#: deployment falls back to when it never pinned one.
DEFAULT_LEGACY = MPNET.name
def _index() -> Dict[str, EmbeddingModel]:
"""Build the lookup table of every accepted spelling."""
table: Dict[str, EmbeddingModel] = {}
for model in MODELS:
for key in (model.name, *model.aliases):
table[key.lower()] = model
return table
_LOOKUP = _index()
def resolve(name: Optional[str]) -> Optional[EmbeddingModel]:
"""Return the registry entry for ``name``, or ``None`` when unknown.
Args:
name: A configured ``EMBEDDINGS_NAME`` value, in any accepted spelling.
Returns:
The matching :class:`EmbeddingModel`, or ``None`` for a name the
registry does not describe -- which callers treat as a Hugging Face
repository rather than an error.
"""
if not name:
return None
return _LOOKUP.get(name.strip().lower())
def dimension_for(name: Optional[str]) -> Optional[int]:
"""Vector width for ``name``, or ``None`` when unknown."""
model = resolve(name)
return model.dimension if model else None
def max_input_tokens_for(name: Optional[str]) -> Optional[int]:
"""Context window for ``name``, or ``None`` when unknown."""
model = resolve(name)
return model.max_input_tokens if model else None
def known_names() -> Tuple[str, ...]:
"""Canonical names of every registered model, for error messages."""
return tuple(model.name for model in MODELS)
+37
View File
@@ -55,6 +55,10 @@ from application.storage.db.repositories.wiki_pages import (
from application.storage.db.session import db_readonly, db_session
from application.storage.db.source_config import SourceConfig
from application.storage.storage_creator import StorageCreator
from application.upload_limits import (
enforce_parseable_attachment,
UnsupportedUploadTypeError,
)
from application.utils import (
count_tokens_docs,
get_encoding,
@@ -1554,6 +1558,30 @@ class AttachmentRejectedError(Exception):
"""
def _reject_unparseable_attachment(
local_path: str, filename: str, parser_extensions
) -> None:
"""Reject an attachment with no parser whose contents are binary.
Defence in depth behind the route's gate, and stricter than it: this runs
against the parser table actually loaded, so a suffix the route trusted
(.webp with docling absent, say) is still refused rather than opened as
plain text by ``SimpleDirectoryReader`` and stored as garbage.
Args:
local_path: Filesystem path of the attachment about to be parsed.
filename: The upload's original filename, which carries the suffix.
parser_extensions: Suffixes the live ``file_extractor`` can parse.
Raises:
AttachmentRejectedError: If the file has no parser and is not text.
"""
try:
enforce_parseable_attachment(local_path, filename, parser_extensions)
except UnsupportedUploadTypeError as exc:
raise AttachmentRejectedError(str(exc)) from exc
def _reject_attachment_zip_bomb(local_path: str) -> None:
"""Reject a zip-container attachment that decompresses to too much.
@@ -1752,6 +1780,7 @@ def attachment_worker(self, file_info, user):
parser_name = type(_parser).__name__ if _parser is not None else "SimpleDirectoryReader"
def _parse_local_file(local_path: str, **kwargs) -> Document:
_reject_unparseable_attachment(local_path, filename, set(file_extractor))
_reject_attachment_zip_bomb(local_path)
parse_path, is_temp_copy = _bounded_attachment_copy(local_path)
try:
@@ -2524,6 +2553,14 @@ def reembed_wiki_page_worker(self, source_id, path, content_hash, user):
with db_session() as conn:
WikiPagesRepository(conn).set_embed_status(source_id, path, "embedded")
# These chunks were just embedded with the configured model, so the
# source now names it. Wiki sources created before this was recorded
# carry NULL, which the boot mismatch check reads as the legacy
# model and reports as stale; stamping here heals them on the next
# page edit.
SourcesRepository(conn).update(
source_id, user, {"model": settings.EMBEDDINGS_NAME}
)
except Exception:
with db_session() as conn:
WikiPagesRepository(conn).set_embed_status(source_id, path, "failed")
+1 -1
View File
@@ -45,7 +45,7 @@ services:
# must set INSTALL_TESSERACT=true before rebuilding (see docker-compose.yaml).
INSTALL_TESSERACT: ${INSTALL_TESSERACT:-false}
# `parsing` queue carries read_document/parse_document; required for its await to resolve.
command: celery -A application.app.celery worker -l INFO -Q docsgpt,parsing
command: celery -A application.app.celery worker -l INFO -Q docsgpt,parsing,embeddings
env_file:
- ../.env
environment:
+1 -1
View File
@@ -40,7 +40,7 @@ services:
user: root
image: arc53/docsgpt:develop
# `parsing` queue carries read_document/parse_document; required for its await to resolve.
command: celery -A application.app.celery worker -l INFO -B -Q docsgpt,parsing
command: celery -A application.app.celery worker -l INFO -B -Q docsgpt,parsing,embeddings
env_file:
- ../.env
environment:
+7 -4
View File
@@ -55,10 +55,13 @@ services:
args:
INSTALL_DOCLING: ${INSTALL_DOCLING:-false}
INSTALL_TESSERACT: ${INSTALL_TESSERACT:-false}
# Consumes the default queue AND the dedicated `parsing` queue (read_document /
# parse_document). Without `parsing` here the read_document await never resolves.
# For heavy/OCR parsing run a separate worker with `-Q parsing`.
command: celery -A application.app.celery worker -l INFO -B -Q docsgpt,parsing
# Consumes the default queue AND the dedicated `parsing` (read_document /
# parse_document) and `embeddings` (query embedding) queues. Without `parsing`
# the read_document await never resolves; without `embeddings` every search
# fails after EMBEDDINGS_DELEGATE_TIMEOUT, because EMBEDDINGS_DELEGATE_TO_WORKER
# is on by default. For heavy/OCR parsing run a separate worker with `-Q parsing`;
# to keep query latency off the ingest pool, another with `-Q embeddings`.
command: celery -A application.app.celery worker -l INFO -B -Q docsgpt,parsing,embeddings
env_file:
- ../.env
environment:
@@ -87,7 +87,7 @@ spec:
image: arc53/docsgpt
# `parsing` queue carries read_document/parse_document; required for its await to resolve.
# For heavy/OCR parsing, run a separate deployment with `-Q parsing` (and GPU env).
command: ["celery", "-A", "application.app.celery", "worker", "-l", "INFO", "-n", "worker.%h", "-Q", "docsgpt,parsing"]
command: ["celery", "-A", "application.app.celery", "worker", "-l", "INFO", "-n", "worker.%h", "-Q", "docsgpt,parsing,embeddings"]
resources:
limits:
memory: "4Gi"
+1 -1
View File
@@ -214,7 +214,7 @@ and leaves this worker light.
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
@@ -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:
+8 -4
View File
@@ -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:** leave it unset and DocsGPT picks for you at first boot, recording the choice so it never changes underneath you: a fresh install is pinned to `ibm-granite/granite-embedding-311m-multilingual-r2` (multilingual, 32k context, same 768 dimensions), and an install that already has sources is pinned to `huggingface_sentence-transformers/all-mpnet-base-v2` so its index stays readable. Setting it here overrides that pin.
- **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.
@@ -468,10 +469,13 @@ 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). |
| `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
+72 -21
View File
@@ -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
@@ -89,7 +83,48 @@ Some remote servers (notably llama.cpp) reject any single input larger than thei
EMBEDDINGS_MAX_INPUT_TOKENS=512
```
When set, each input string is truncated to that many tokens and the overflow is dropped (lossy by design). Token counts use DocsGPT's shared tiktoken encoding, which differs from your server's tokenizer, so choose a limit with some headroom below the server's true limit to absorb tokenizer skew. Leave the setting unset (or `0`) to disable truncation.
When set, each input string is truncated to that many tokens and the overflow is dropped (lossy by design).
You usually do not need to set it. When `EMBEDDINGS_NAME` names a model DocsGPT knows, its context window is used automatically, and counting uses that model's own tokenizer, so the limit and the count are in the same unit. Set `EMBEDDINGS_MAX_INPUT_TOKENS` when your server serves a model DocsGPT does not know, or when you deliberately want a limit below the model's own.
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. **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
```
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
@@ -97,8 +132,24 @@ Each embedding model produces vectors of a fixed dimension, and your vector stor
If you need to switch embedding models, you must re-ingest your sources so the index is rebuilt with the new dimension. This also applies to the [GraphRAG](/Sources/GraphRAG) graph tables, which are sized to the embedding dimension at creation time.
### A matching dimension is not a matching model
The dimension check is a guard against a corrupt index, not a guarantee that a swap is safe. Two models of the *same* width — `all-mpnet-base-v2` and `granite-embedding-311m-multilingual-r2` are both 768 — raise nothing at all, and every query is then embedded by a different model than the stored vectors were. Nothing fails; retrieval quality simply degrades.
Switching between same-width models therefore still requires re-embedding:
```bash
python -m application.scripts.reembed
```
Run it after changing `EMBEDDINGS_NAME` and before serving queries. See [Upgrading](/upgrading) for the granite migration specifically.
Changing to a model of a *different* width is supported for FAISS: the script rebuilds the index at the new width and keeps the existing chunk ids. For `pgvector` the vector column is sized at creation time, so a width change there still means re-ingesting those sources.
With `GRAPHRAG_ENABLED`, the script also rewrites `graph_nodes.name_embedding` on `pgvector`. Those vectors seed every graph traversal, and they share the chunk vectors' width, so leaving them behind degrades graph retrieval just as silently.
## 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.
+79
View File
@@ -11,6 +11,85 @@ import { Callout } from 'nextra/components'
**Upgrading from 0.16.x?** User data moved from MongoDB to Postgres in 0.17.0. Follow the [Postgres Migration guide](/Deploying/Postgres-Migration) before running `docker compose pull` or `git pull` — existing deployments will not start cleanly without it.
</Callout>
## 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 **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.
### Switching an existing deployment to granite
<Callout type="error">
Changing `EMBEDDINGS_NAME` on an index that already has vectors **breaks retrieval silently**. Both models are 768-dimensional, so nothing raises an error — queries are simply compared against vectors that mean something else, and answers quietly get worse. Always re-embed.
</Callout>
Set the model, then rebuild the vectors:
```bash
# 1. In your .env
EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2
# 2. Rebuild the vectors from the chunk text already in your index
docker compose exec backend python -m application.scripts.reembed --dry-run
docker compose exec backend python -m application.scripts.reembed
```
Re-embedding reads the chunk text already stored in your index. It does not re-download, re-parse or re-chunk your documents, so no source files are needed and the run is proportional to index size, not corpus size. Both `pgvector` and `faiss` are supported.
Useful flags:
| Flag | Effect |
| --- | --- |
| `--dry-run` | Report how many chunks would change, write nothing |
| `--sources a,b` | Only these source ids — also how you retry a failed source |
| `--batch-size N` | Chunks per embed call (default 64) |
The script processes sources independently: one failing source is reported and skipped rather than aborting the run, and the exit code is non-zero if any failed. For `pgvector` it reads a page of chunks at a time and updates rows in place, so memory stays flat on a large index and an interrupted run simply re-does its last batch. For `faiss` it builds the replacement index in memory, writes each file to a temporary path, and moves it into place — so an interrupt during either the rebuild or the write leaves the existing index intact rather than truncated.
<Callout type="warning">
Stop ingest before you run this. It reads each source's chunks and writes the vectors back; anything ingested while it runs can be overwritten by the rebuild (`faiss`) or missed by it (`pgvector`).
</Callout>
<Callout type="info">
Running [GraphRAG](/Sources/GraphRAG)? The script also rewrites `graph_nodes.name_embedding`, which seeds every graph traversal. Those vectors are written once at extraction time and share the chunk vectors' width, so leaving them in the old model's space degrades graph retrieval just as silently as the chunk vectors would — and needs no LLM re-extraction to fix.
</Callout>
### Custom local models
A local model now runs through ONNX Runtime, so its repository must ship an ONNX
export (`onnx/model.onnx`) or be one of FastEmbed's built-in models. Repositories
with PyTorch weights only no longer load; `hkunlp/instructor-large`, previously
supported by name, is one of them. Serve such a model over `EMBEDDINGS_BASE_URL`
instead, or switch to a model with an export.
How to run the model — pooling, and whether outputs are L2-normalised — is read
from the repository's own `1_Pooling/config.json` and `modules.json`. Two cases
need attention:
- **Models with a Dense projection layer** (`sentence-transformers/LaBSE`,
`distiluse-base-multilingual-cased-v1`) are now **refused at startup**.
FastEmbed cannot apply the projection, so it would have produced vectors of the
wrong width in a different space. If you were running one, its stored vectors
were already wrong; move it to `EMBEDDINGS_BASE_URL` or pick another model.
- **Repositories that declare nothing** fall back to mean pooling with
normalisation and log a warning. Pin the real values with `EMBEDDINGS_POOLING`
(`cls` or `mean`) and `EMBEDDINGS_NORMALIZE`.
<Callout type="info">
Staying on `all-mpnet-base-v2` is a supported choice — it remains in the model registry and in `setup.sh`. You only need this section if you want to move to granite.
</Callout>
## Check your version
```bash
File diff suppressed because it is too large. Load diff
@@ -11,6 +11,7 @@ interface FetchAnswerStreamingProps {
conversationId?: string | null;
apiHost?: string;
onEvent?: (event: MessageEvent) => void;
signal?: AbortSignal;
}
export interface FeedbackPayload {
@@ -31,13 +32,17 @@ export function fetchAnswerStreaming({
onEvent = () => {
console.log('Event triggered, but no handler provided.');
},
signal,
}: FetchAnswerStreamingProps): Promise<void> {
return new Promise<void>((resolve, reject) => {
const body = {
question: question,
history: JSON.stringify(history),
history: JSON.stringify(
history
.filter((item) => item.prompt && item.response)
.map(({ prompt, response }) => ({ prompt, response })),
),
conversation_id: conversationId,
model: 'default',
api_key: apiKey,
};
fetch(apiHost + '/stream', {
@@ -46,49 +51,61 @@ export function fetchAnswerStreaming({
'Content-Type': 'application/json',
},
body: JSON.stringify(body),
signal,
})
.then((response) => {
if (!response.body) throw Error('No response body');
const reader = response.body.getReader();
const decoder = new TextDecoder('utf-8');
let counter = 0; // eslint-disable-line @typescript-eslint/no-unused-vars
let buffer = '';
const emit = (rawLine: string) => {
const line = rawLine.trim();
if (line === '') return;
if (!line.startsWith('data:')) return;
const payload = line.substring(5).trim();
if (payload === '' || payload === '[DONE]') return;
onEvent(new MessageEvent('message', { data: payload }));
};
const processStream = ({
done,
value,
}: ReadableStreamReadResult<Uint8Array>) => {
if (done) {
if (buffer.trim() !== '') emit(buffer);
buffer = '';
resolve();
return;
}
buffer += decoder.decode(value, { stream: true });
counter += 1;
const lines = buffer.split('\n');
buffer = lines.pop() ?? '';
for (const line of lines) emit(line);
const chunk = decoder.decode(value);
const lines = chunk.split('\n');
for (let line of lines) {
if (line.trim() == '') {
continue;
}
if (line.startsWith('data:')) {
line = line.substring(5);
}
const messageEvent = new MessageEvent('message', {
data: line,
});
onEvent(messageEvent); // handle each message
}
reader.read().then(processStream).catch(reject);
reader.read().then(processStream).catch(onReadError);
};
reader.read().then(processStream).catch(reject);
const onReadError = (error: unknown) => {
if (signal?.aborted || (error as Error)?.name === 'AbortError') {
resolve();
return;
}
reject(error);
};
reader.read().then(processStream).catch(onReadError);
})
.catch((error) => {
if (signal?.aborted || (error as Error)?.name === 'AbortError') {
resolve();
return;
}
console.error('Connection failed:', error);
reject(error);
});
@@ -14,6 +14,11 @@ declare module 'styled-components' {
};
/** Present only in SearchBar theme */
name?: string;
/** Gradient stops for the swept status text. */
shimmer?: {
base: string;
highlight: string;
};
/** Present only in DocsGPTWidget theme (always provided when these styled components render) */
dimensions?: {
size: string;
@@ -41,6 +46,12 @@ export interface Query {
sources?: { title: string; text: string; source: string }[];
conversationId?: string | null;
title?: string | null;
/** Accumulated from thought events. */
thought?: string;
/** Latest notice or running workflow node; drives the status line. */
notice?: string;
/** Tool names from tool_calls / tool_call. */
toolCalls?: string[];
}
export interface WidgetProps {
@@ -0,0 +1,50 @@
export interface StreamEvent {
type?: string;
[key: string]: unknown;
}
export interface ToolCallEvent {
action_name?: string;
tool_name?: string;
}
/** internal_search -> Internal Search */
export const prettifyName = (name: string): string =>
name
.split(/[_\s]+/)
.filter(Boolean)
.map((word) => word.charAt(0).toUpperCase() + word.slice(1).toLowerCase())
.join(' ');
/**
* Status-line text for a running workflow node, or null when the node is
* scaffolding the user gains nothing from seeing. Keyed on node_type (a
* closed enum) rather than node_title, which is free text like "agent node".
*/
export const workflowStepLabel = (event: StreamEvent): string | null => {
if (event.status !== 'running') return null;
switch (event.node_type) {
case 'agent':
return 'Thinking…';
case 'code':
return 'Running code…';
case 'condition':
return 'Deciding next step…';
case 'start':
case 'end':
case 'state':
case 'note':
return null;
default:
return typeof event.node_title === 'string' && event.node_title
? `${prettifyName(event.node_title)}…`
: null;
}
};
export const toolNames = (calls: unknown): string[] => {
if (!Array.isArray(calls)) return [];
return calls
.map((call: ToolCallEvent | null) => call?.action_name ?? call?.tool_name)
.filter((name): name is string => Boolean(name));
};
+65 -10
View File
@@ -48,7 +48,10 @@ import { useArmedSend } from './message-input/armedSend';
import { handleAbort } from '../conversation/conversationSlice';
import {
AUDIO_FILE_ACCEPT_ATTR,
FILE_UPLOAD_ACCEPT,
getFileExtension,
parseUploadErrorMessage,
parseUploadErrorsByIndex,
partitionAttachmentFiles,
} from '../constants/fileUpload';
import { UserToolType } from '../settings/types';
import { isChatToolVisible } from '../utils/toolUtils';
@@ -497,8 +500,31 @@ export default function MessageInput({
);
const uploadFiles = useCallback(
(files: File[]) => {
if (!files || files.length === 0) return;
async (incomingFiles: File[]) => {
if (!incomingFiles || incomingFiles.length === 0) return;
// Run the server's own rule here, not just the input's `accept`:
// mobile pickers ignore `accept`, and a file the server will refuse
// should say so before it costs an upload. Surface the refusal as a
// failed chip so the user sees why instead of a silent drop.
const { supported, unsupported } =
await partitionAttachmentFiles(incomingFiles);
unsupported.forEach((file) => {
dispatch(
addAttachment({
id: generateId(),
fileName: file.name,
progress: 0,
status: 'failed' as const,
taskId: '',
errorMessage: t('conversation.attachments.unsupportedType', {
extension: getFileExtension(file.name) || '?',
}),
}),
);
});
if (supported.length === 0) return;
const files = supported;
const apiHost = import.meta.env.VITE_API_HOST;
@@ -579,9 +605,14 @@ export default function MessageInput({
tasksByIndex.set(uploadIndex, task);
});
const errorsByIndex = new Map<number, string | undefined>();
errors.forEach((errorItem) => {
if (typeof errorItem.upload_index === 'number') {
failedIndices.add(errorItem.upload_index);
errorsByIndex.set(
errorItem.upload_index,
errorItem.error,
);
}
});
@@ -616,7 +647,10 @@ export default function MessageInput({
dispatch(
updateAttachment({
id: uiId,
updates: { status: 'failed' },
updates: {
status: 'failed',
errorMessage: errorsByIndex.get(index),
},
}),
);
return;
@@ -748,11 +782,20 @@ export default function MessageInput({
}
} else {
console.error('Upload failed', status, xhr.responseText);
Object.values(indexToUiId).forEach((id) =>
// Each file gets its own reason where the server sent one; the
// top-level message is the fallback, not the answer for all of
// them — a batch can fail two files for two different reasons.
const fallbackMessage = parseUploadErrorMessage(xhr.responseText);
const errorsByIndex = parseUploadErrorsByIndex(xhr.responseText);
Object.entries(indexToUiId).forEach(([index, id]) =>
dispatch(
updateAttachment({
id,
updates: { status: 'failed' },
updates: {
status: 'failed',
errorMessage:
errorsByIndex.get(Number(index)) ?? fallbackMessage,
},
}),
),
);
@@ -871,7 +914,10 @@ export default function MessageInput({
dispatch(
updateAttachment({
id: uniqueId,
updates: { status: 'failed' },
updates: {
status: 'failed',
errorMessage: parseUploadErrorMessage(xhr.responseText),
},
}),
);
}
@@ -891,7 +937,7 @@ export default function MessageInput({
xhr.send(formData);
});
},
[dispatch, token, trackAttachment],
[dispatch, t, token, trackAttachment],
);
const handleFileAttachment = (e: React.ChangeEvent<HTMLInputElement>) => {
@@ -928,7 +974,10 @@ export default function MessageInput({
setHandleDragActive(false);
},
maxSize: 25000000,
accept: FILE_UPLOAD_ACCEPT,
// No `accept`: react-dropzone would drop a rejected file on the floor
// with no feedback, and its mime matching disagrees with the server for
// text files that have no parser (.py, .log). uploadFiles applies the
// server's rule and reports what it refuses.
});
const handleInput = useCallback(() => {
@@ -1594,7 +1643,13 @@ export default function MessageInput({
onChange={handleVoiceFileAttachment}
/>
<div className="border-border bg-card relative flex w-full flex-col rounded-3xl border dark:bg-transparent">
{/* translate="no": keep Chrome's page translator out of the composer —
it rewrites text nodes into <font> wrappers and React loses the
controls (dead Attach button, see AttachFileButton). */}
<div
translate="no"
className="border-border bg-card relative flex w-full flex-col rounded-3xl border dark:bg-transparent"
>
<AttachmentChipList
attachments={attachments}
draggingId={draggingId}
@@ -1,7 +1,7 @@
import { useTranslation } from 'react-i18next';
import ClipIcon from '../../assets/clip.svg';
import { FILE_UPLOAD_ACCEPT_ATTR } from '../../constants/fileUpload';
import { ATTACHMENT_FILE_ACCEPT_ATTR } from '../../constants/fileUpload';
type AttachFileButtonProps = {
onChange: (e: React.ChangeEvent<HTMLInputElement>) => void;
@@ -11,7 +11,14 @@ export default function AttachFileButton({ onChange }: AttachFileButtonProps) {
const { t } = useTranslation();
return (
<label className="xs:px-3 xs:py-1.5 dark:border-border border-border hover:bg-muted dark:hover:bg-muted flex cursor-pointer items-center rounded-full border px-2 py-1 transition-colors">
// translate="no": Chrome's page translator rewrites text nodes inside
// this label into <font> wrappers and React then loses the control — a
// user with auto-translate on rage-clicked a dead Attach button until
// they reloaded. Keep the composer's controls out of the translator.
<label
translate="no"
className="xs:px-3 xs:py-1.5 dark:border-border border-border hover:bg-muted dark:hover:bg-muted flex cursor-pointer items-center rounded-full border px-2 py-1 transition-colors"
>
<img
src={ClipIcon}
alt="Attach"
@@ -24,7 +31,7 @@ export default function AttachFileButton({ onChange }: AttachFileButtonProps) {
type="file"
className="hidden"
multiple
accept={FILE_UPLOAD_ACCEPT_ATTR}
accept={ATTACHMENT_FILE_ACCEPT_ATTR}
onChange={onChange}
/>
</label>
@@ -25,94 +25,122 @@ export default function AttachmentChipList({
}: AttachmentChipListProps) {
const { t } = useTranslation();
// A tooltip is the one place a touch user can never look, and this list is
// where a phone picker's unsupported file lands. Show the reason inline,
// as soon as it is known, rather than only once a send is attempted.
const failures = attachments.filter(
(attachment) => attachment.status === 'failed' && attachment.errorMessage,
);
return (
<div className="flex flex-wrap gap-1.5 px-2 py-2 sm:gap-2 sm:px-3">
{attachments.map((attachment) => {
return (
<div
key={attachment.id}
draggable={true}
onDragStart={(e) => onDragStart(e, attachment.id)}
onDragOver={onDragOver}
onDrop={(e) => onDropOn(e, attachment.id)}
className={`group dark:text-foreground bg-muted text-muted-foreground dark:bg-accent relative flex items-center rounded-xl px-2 py-1 text-xs sm:px-3 sm:py-1.5 sm:text-sm ${
attachment.status !== 'completed' ? 'opacity-70' : 'opacity-100'
} ${
draggingId === attachment.id
? 'ring-dashed opacity-60 ring-2 ring-purple-200'
: ''
}`}
title={attachment.fileName}
>
<div className="bg-primary mr-2 flex h-8 w-8 items-center justify-center rounded-md p-1">
{attachment.status === 'completed' && (
<img
src={DocumentationDark}
alt="Attachment"
className="h-[15px] w-[15px] object-fill"
/>
)}
{attachment.status === 'failed' && (
<img
src={AlertIcon}
alt="Failed"
className="h-[15px] w-[15px] object-fill"
/>
)}
{(attachment.status === 'uploading' ||
attachment.status === 'processing') && (
<div className="flex h-[15px] w-[15px] items-center justify-center">
<svg className="h-[15px] w-[15px]" viewBox="0 0 24 24">
<circle
className="opacity-0"
cx="12"
cy="12"
r="10"
stroke="transparent"
strokeWidth="4"
fill="none"
/>
<circle
className="text-[#ECECF1]"
cx="12"
cy="12"
r="10"
stroke="currentColor"
strokeWidth="4"
fill="none"
strokeDasharray="62.83"
strokeDashoffset={62.83 * (1 - attachment.progress / 100)}
transform="rotate(-90 12 12)"
/>
</svg>
</div>
)}
</div>
<span className="max-w-[120px] truncate font-medium sm:max-w-[150px]">
{attachment.fileName}
</span>
<Button
type="button"
variant="ghost"
size="icon-sm"
className="ml-1.5 h-auto w-auto rounded-full p-1"
onClick={() => {
onRemove(attachment.id);
}}
aria-label={t('conversation.attachments.remove')}
<>
<div className="flex flex-wrap gap-1.5 px-2 py-2 sm:gap-2 sm:px-3">
{attachments.map((attachment) => {
return (
<div
key={attachment.id}
draggable={true}
onDragStart={(e) => onDragStart(e, attachment.id)}
onDragOver={onDragOver}
onDrop={(e) => onDropOn(e, attachment.id)}
className={`group dark:text-foreground bg-muted text-muted-foreground dark:bg-accent relative flex items-center rounded-xl px-2 py-1 text-xs sm:px-3 sm:py-1.5 sm:text-sm ${
attachment.status !== 'completed' ? 'opacity-70' : 'opacity-100'
} ${
draggingId === attachment.id
? 'ring-dashed opacity-60 ring-2 ring-purple-200'
: ''
}`}
title={
attachment.status === 'failed' && attachment.errorMessage
? `${attachment.fileName}: ${attachment.errorMessage}`
: attachment.fileName
}
>
<X
<div className="bg-primary mr-2 flex h-8 w-8 items-center justify-center rounded-md p-1">
{attachment.status === 'completed' && (
<img
src={DocumentationDark}
alt="Attachment"
className="h-[15px] w-[15px] object-fill"
/>
)}
{attachment.status === 'failed' && (
<img
src={AlertIcon}
alt="Failed"
className="h-[15px] w-[15px] object-fill"
/>
)}
{(attachment.status === 'uploading' ||
attachment.status === 'processing') && (
<div className="flex h-[15px] w-[15px] items-center justify-center">
<svg className="h-[15px] w-[15px]" viewBox="0 0 24 24">
<circle
className="opacity-0"
cx="12"
cy="12"
r="10"
stroke="transparent"
strokeWidth="4"
fill="none"
/>
<circle
className="text-[#ECECF1]"
cx="12"
cy="12"
r="10"
stroke="currentColor"
strokeWidth="4"
fill="none"
strokeDasharray="62.83"
strokeDashoffset={
62.83 * (1 - attachment.progress / 100)
}
transform="rotate(-90 12 12)"
/>
</svg>
</div>
)}
</div>
<span className="max-w-[120px] truncate font-medium sm:max-w-[150px]">
{attachment.fileName}
</span>
<Button
type="button"
variant="ghost"
size="icon-sm"
className="ml-1.5 h-auto w-auto rounded-full p-1"
onClick={() => {
onRemove(attachment.id);
}}
aria-label={t('conversation.attachments.remove')}
className="h-2.5 w-2.5"
/>
</Button>
</div>
);
})}
</div>
>
<X
aria-label={t('conversation.attachments.remove')}
className="h-2.5 w-2.5"
/>
</Button>
</div>
);
})}
</div>
{failures.length > 0 && (
<div
className="flex flex-col gap-0.5 px-2 pb-1 text-xs text-[#B42318] sm:px-3"
role="alert"
>
{failures.map((attachment) => (
<span key={attachment.id}>
{attachment.fileName}: {attachment.errorMessage}
</span>
))}
</div>
)}
</>
);
}
+240
View File
@@ -0,0 +1,240 @@
import { describe, expect, it } from 'vitest';
import {
ATTACHMENT_FILE_ACCEPT_ATTR,
ATTACHMENT_PARSER_EXTENSIONS,
getFileExtension,
hasAttachmentParser,
looksLikeText,
parseUploadErrorMessage,
parseUploadErrorsByIndex,
partitionAttachmentFiles,
} from './fileUpload';
const file = (name: string, body: BlobPart = 'x', type = '') =>
new File([body], name, { type });
const bytes = (...values: number[]) => new Uint8Array(values);
const MP4_HEADER = bytes(
0x00,
0x00,
0x00,
0x18,
0x66,
0x74,
0x79,
0x70,
0x69,
0x73,
0x6f,
0x6d,
);
describe('attachment type gate', () => {
it('extracts a lower-cased extension', () => {
expect(getFileExtension('Photo.JPG')).toBe('.jpg');
expect(getFileExtension('a.tar.gz')).toBe('.gz');
expect(getFileExtension('noext')).toBe('');
expect(getFileExtension('.env')).toBe('');
});
it('recognises parser-backed suffixes by name, case-insensitively', () => {
expect(hasAttachmentParser(file('Report.PDF'))).toBe(true);
expect(hasAttachmentParser(file('scan.WebP'))).toBe(true);
expect(hasAttachmentParser(file('subs.vtt'))).toBe(true);
// No parser — these are judged on content, not on the suffix.
expect(hasAttachmentParser(file('main.py'))).toBe(false);
expect(hasAttachmentParser(file('archive.zip'))).toBe(false);
expect(hasAttachmentParser(file('clip.mp4'))).toBe(false);
// .txt is the plain-text fallthrough itself, not a parser.
expect(hasAttachmentParser(file('notes.txt'))).toBe(false);
});
it('mirrors the backend list: no zip, no video, images and markup included', () => {
expect(ATTACHMENT_PARSER_EXTENSIONS).toContain('.pdf');
// Parser-backed suffixes the first cut of this list missed.
for (const ext of [
'.webp',
'.tiff',
'.tif',
'.bmp',
'.vtt',
'.xml',
'.mdx',
]) {
expect(ATTACHMENT_PARSER_EXTENSIONS).toContain(ext);
}
expect(ATTACHMENT_PARSER_EXTENSIONS).not.toContain('.zip');
expect(ATTACHMENT_PARSER_EXTENSIONS).not.toContain('.mp4');
expect(ATTACHMENT_PARSER_EXTENSIONS).not.toContain('.txt');
});
it('never hides a file the upload would accept behind the picker filter', () => {
const listed = ATTACHMENT_FILE_ACCEPT_ATTR.split(',');
for (const ext of ATTACHMENT_PARSER_EXTENSIONS) {
expect(listed).toContain(ext);
}
// Parserless text (.txt, .py, .log) is accepted on content, so the
// picker must not filter it out.
expect(listed).toContain('.txt');
expect(listed).toContain('text/*');
});
});
describe('looksLikeText', () => {
it('accepts empty, plain, accented and ANSI-coloured text', () => {
expect(looksLikeText(new Uint8Array())).toBe(true);
expect(looksLikeText(new TextEncoder().encode('plain text\n'))).toBe(true);
expect(looksLikeText(new TextEncoder().encode('café — em dash\n'))).toBe(
true,
);
expect(looksLikeText(bytes(0x1b, 0x5b, 0x33, 0x31, 0x6d, 0x6f, 0x6b))).toBe(
true,
);
});
it('accepts BOM-marked Unicode, which is half NUL bytes but still text', () => {
// UTF-8, UTF-16 LE/BE, and UTF-32 BE (no TextDecoder — left to the server).
expect(looksLikeText(bytes(0xef, 0xbb, 0xbf, 0x68, 0x69))).toBe(true);
expect(looksLikeText(bytes(0xff, 0xfe, 0x68, 0x00, 0x69, 0x00))).toBe(true);
expect(looksLikeText(bytes(0xfe, 0xff, 0x00, 0x68, 0x00, 0x69))).toBe(true);
expect(looksLikeText(bytes(0x00, 0x00, 0xfe, 0xff, 0x00, 0x68))).toBe(true);
});
it('still rejects binary behind a BOM — three bytes buy nothing', () => {
// UTF-8 BOM + MP4 header: the byte rules apply to what follows it.
expect(looksLikeText(bytes(0xef, 0xbb, 0xbf, ...MP4_HEADER))).toBe(false);
// UTF-16 LE BOM + control-only code points.
expect(
looksLikeText(
bytes(0xff, 0xfe, 0x01, 0x00, 0x02, 0x00, 0x03, 0x00, 0x04, 0x00),
),
).toBe(false);
// UTF-16 BE BOM + bytes that decode to NUL characters.
expect(looksLikeText(bytes(0xfe, 0xff, 0x00, 0x00, 0x00, 0x41))).toBe(
false,
);
});
it('rejects a NUL byte and dense control bytes', () => {
expect(looksLikeText(MP4_HEADER)).toBe(false);
expect(looksLikeText(bytes(0x74, 0x78, 0x00, 0x74))).toBe(false);
expect(
looksLikeText(
new Uint8Array(Array.from({ length: 64 }, (_, i) => i + 1)),
),
).toBe(false);
});
});
describe('partitionAttachmentFiles', () => {
it('refuses video and archives whatever the picker claims the type is', async () => {
const { supported, unsupported } = await partitionAttachmentFiles([
file('clip.mp4', MP4_HEADER, 'text/plain'),
file('archive.zip', bytes(0x50, 0x4b, 0x03, 0x04, 0x00, 0x00)),
]);
expect(supported).toEqual([]);
expect(unsupported.map((f) => f.name)).toEqual(['clip.mp4', 'archive.zip']);
});
it('keeps text files that have no parser — the backend reads them', async () => {
const { supported, unsupported } = await partitionAttachmentFiles([
file('notes.txt', 'plain\n'),
file('main.py', 'def main():\n return 1\n'),
file('server.log', '2026-09-02 ERROR boom\n'),
file('Dockerfile', 'FROM python:3.12\n'),
]);
expect(supported.map((f) => f.name)).toEqual([
'notes.txt',
'main.py',
'server.log',
'Dockerfile',
]);
expect(unsupported).toEqual([]);
});
it('refuses a binary renamed to .txt — .txt is sniffed like any other suffix', async () => {
const { supported, unsupported } = await partitionAttachmentFiles([
file('notes.txt', MP4_HEADER),
// A BOM in front of it changes nothing.
file('bom.txt', bytes(0xef, 0xbb, 0xbf, ...MP4_HEADER)),
]);
expect(supported).toEqual([]);
expect(unsupported.map((f) => f.name)).toEqual(['notes.txt', 'bom.txt']);
});
it('admits parser-backed types on their suffix, binary contents and all', async () => {
// A PDF is binary; the sniff must never be applied to it.
const { supported, unsupported } = await partitionAttachmentFiles([
file('paper.pdf', MP4_HEADER),
file('photo.jpeg', MP4_HEADER),
]);
expect(supported.map((f) => f.name)).toEqual(['paper.pdf', 'photo.jpeg']);
expect(unsupported).toEqual([]);
});
it('partitions a mixed selection and keeps the original order', async () => {
const mp4 = file('clip.mp4', MP4_HEADER);
const txt = file('notes.txt');
const pdf = file('paper.pdf');
const { supported, unsupported } = await partitionAttachmentFiles([
mp4,
txt,
pdf,
]);
expect(supported).toEqual([txt, pdf]);
expect(unsupported).toEqual([mp4]);
});
});
describe('parseUploadErrorMessage', () => {
it('returns the server message when present', () => {
expect(
parseUploadErrorMessage(
'{"success":false,"message":"Unsupported file type: .mp4"}',
),
).toBe('Unsupported file type: .mp4');
});
it('returns undefined for non-JSON bodies or a missing message', () => {
expect(parseUploadErrorMessage('<html>502</html>')).toBeUndefined();
expect(parseUploadErrorMessage('{"ok":1}')).toBeUndefined();
expect(parseUploadErrorMessage('')).toBeUndefined();
});
});
describe('parseUploadErrorsByIndex', () => {
it('keys each reason by its upload_index', () => {
const byIndex = parseUploadErrorsByIndex(
JSON.stringify({
success: false,
message: 'Unsupported file type: .mp4',
errors: [
{
upload_index: 0,
filename: 'clip.mp4',
error: 'Unsupported file type: .mp4',
},
{
upload_index: 1,
filename: 'a.zip',
error: 'Unsupported file type: .zip',
},
],
}),
);
// Without this, both chips would repeat the first file's reason.
expect(byIndex.get(0)).toBe('Unsupported file type: .mp4');
expect(byIndex.get(1)).toBe('Unsupported file type: .zip');
});
it('is empty for bodies with no usable errors array', () => {
expect(parseUploadErrorsByIndex('<html>502</html>').size).toBe(0);
expect(parseUploadErrorsByIndex('{"message":"nope"}').size).toBe(0);
expect(parseUploadErrorsByIndex('').size).toBe(0);
expect(
parseUploadErrorsByIndex('{"errors":[{"error":"no index"}]}').size,
).toBe(0);
});
});
+243
View File
@@ -129,3 +129,246 @@ export const SOURCE_FILE_TREE_ACCEPT_ATTR = [
'.jpg',
'.jpeg',
].join(',');
/**
* Chat-attachment suffixes with a dedicated parser. Mirrors the backend's
* `ATTACHMENT_PARSER_EXTENSIONS` (application/parser/file/constants.py) —
* update both together. Zip is absent: source ingestion extracts archives,
* the attachment path does not.
*
* Not the whole allow-list, and `.txt` is deliberately not here: a suffix
* that isn't listed (.txt, .py, .log, .yaml) is read by the backend's
* plain-text fallthrough, so it is judged on content by
* `partitionAttachmentFiles` instead, exactly as the server does.
*/
export const ATTACHMENT_PARSER_EXTENSIONS: readonly string[] = [
'.rst',
'.md',
'.mdx',
'.pdf',
'.docx',
'.csv',
'.epub',
'.html',
'.xhtml',
'.json',
'.xlsx',
'.pptx',
// Read by the anydoc engine, a core backend dependency: legacy and
// macro-enabled Office, OpenDocument, RTF.
'.docm',
'.doc',
'.odt',
'.rtf',
'.pptm',
'.ppt',
'.pps',
'.ppsx',
'.ppsm',
'.pot',
'.odp',
'.xlsm',
'.xlsb',
'.xls',
'.ods',
'.adoc',
'.asciidoc',
'.png',
'.jpg',
'.jpeg',
'.tiff',
'.tif',
'.bmp',
'.webp',
'.vtt',
'.xml',
'.wav',
'.mp3',
'.m4a',
'.ogg',
'.webm',
];
/**
* Picker filter for the Attach button. A hint only — pickers may ignore it,
* and it must never be narrower than what the upload accepts: `text/*` (plus
* `.txt` explicitly) keeps parserless text files such as .py and .log
* selectable, since the gate reads those happily.
*/
export const ATTACHMENT_FILE_ACCEPT_ATTR = [
...ATTACHMENT_PARSER_EXTENSIONS,
'.txt',
'text/*',
].join(',');
/** Lower-cased last extension including the dot, or '' (dotfiles have none). */
export function getFileExtension(name: string): string {
const base = name.split(/[\\/]/).pop() ?? '';
const dot = base.lastIndexOf('.');
if (dot <= 0) return '';
return base.slice(dot).toLowerCase();
}
/**
* Whether the backend has a parser for this suffix. False is not a refusal:
* the file then has to read as text. Never decided by the mime type the
* picker reports — mobile pickers ignore `accept` and lie about types.
*/
export function hasAttachmentParser(file: { name: string }): boolean {
const ext = getFileExtension(file.name);
return ext !== '' && ATTACHMENT_PARSER_EXTENSIONS.includes(ext);
}
// Mirrors application/upload_limits.py: enough of the head to recognise a
// container header, and a tolerance that keeps real text (UTF-8 accents, an
// ANSI-coloured log) in while a random binary stays out.
const TEXT_SNIFF_BYTES = 8192;
const MAX_NONTEXT_RATIO = 0.1;
// Control bytes that occur in ordinary text: tab, LF, VT, FF, CR, ESC.
const TEXT_CONTROL_BYTES = new Set([0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x1b]);
const TEXT_CONTROL_CHARS = new Set(['\t', '\n', '\v', '\f', '\r', '\u001b']);
// \p{C}: control, format, private-use, surrogate and unassigned code points —
// what binary decodes into.
const NONTEXT_CHAR = /\p{C}/u;
// A UTF-16/32 file is half NUL bytes, so the byte rules below would reject it.
// Its BOM says which encoding to read it as; the decoded characters are then
// judged instead. Longest BOM first — UTF-32-LE starts with the UTF-16-LE one.
// A BOM is never a verdict on its own: three bytes must not buy a video a pass.
const TEXT_BOMS: { bom: number[]; encoding: string | null }[] = [
// TextDecoder has no UTF-32; leave those two to the server, which does.
{ bom: [0xff, 0xfe, 0x00, 0x00], encoding: null },
{ bom: [0x00, 0x00, 0xfe, 0xff], encoding: null },
{ bom: [0xef, 0xbb, 0xbf], encoding: 'utf-8' },
{ bom: [0xff, 0xfe], encoding: 'utf-16le' },
{ bom: [0xfe, 0xff], encoding: 'utf-16be' },
];
/**
* Whether decoded characters read as text rather than decoded binary.
* Replacement characters count against it, so bytes the decoder could not
* read are evidence rather than silently dropped.
*/
function decodedLooksLikeText(text: string): boolean {
// The byte-level NUL rule, one level up: no text holds a NUL character.
if (text.includes('\u0000')) return false;
let total = 0;
let nontext = 0;
for (const char of text) {
total += 1;
if (
char === '\ufffd' ||
(!TEXT_CONTROL_CHARS.has(char) && NONTEXT_CHAR.test(char))
)
nontext += 1;
}
if (total === 0) return true;
return nontext / total <= MAX_NONTEXT_RATIO;
}
/**
* Whether a leading byte sample reads as text rather than binary. A UTF-16
* BOM switches the test to the decoded characters; the content is still what
* decides. Mirrors `looks_like_text` in application/upload_limits.py.
*/
export function looksLikeText(sample: Uint8Array): boolean {
if (sample.length === 0) return true;
let body = sample;
const marked = TEXT_BOMS.find(({ bom }) =>
bom.every((byte, i) => sample[i] === byte),
);
if (marked) {
// UTF-32, or a decoder this browser lacks: defer to the server rather
// than refuse a file it would accept.
if (marked.encoding === null) return true;
body = sample.subarray(marked.bom.length);
if (body.length === 0) return true;
if (marked.encoding !== 'utf-8') {
try {
return decodedLooksLikeText(
new TextDecoder(marked.encoding).decode(body),
);
} catch {
return true;
}
}
// UTF-8 BOM: the byte rules still apply to everything after it.
}
let nontext = 0;
for (const byte of body) {
if (byte === 0x00) return false;
if ((byte < 0x20 && !TEXT_CONTROL_BYTES.has(byte)) || byte === 0x7f)
nontext += 1;
}
return nontext / body.length <= MAX_NONTEXT_RATIO;
}
async function isSupportedAttachmentFile(file: File): Promise<boolean> {
if (hasAttachmentParser(file)) return true;
try {
const head = await file.slice(0, TEXT_SNIFF_BYTES).arrayBuffer();
return looksLikeText(new Uint8Array(head));
} catch {
// Can't read it here — let the upload run and the server decide.
return true;
}
}
/**
* Split a selection into what the backend can read and what it will refuse,
* preserving order. Files with a parser pass on their suffix; the rest are
* sniffed, so source, config and log files go through and a video does not.
*/
export async function partitionAttachmentFiles(files: File[]): Promise<{
supported: File[];
unsupported: File[];
}> {
const verdicts = await Promise.all(files.map(isSupportedAttachmentFile));
const supported: File[] = [];
const unsupported: File[] = [];
files.forEach((file, index) => {
(verdicts[index] ? supported : unsupported).push(file);
});
return { supported, unsupported };
}
/** The ``message`` of a JSON error body, if the server sent one. */
export function parseUploadErrorMessage(body: string): string | undefined {
if (!body) return undefined;
try {
const parsed = JSON.parse(body) as { message?: unknown };
return typeof parsed?.message === 'string' && parsed.message
? parsed.message
: undefined;
} catch {
return undefined;
}
}
/**
* Per-file reasons from an error body, keyed by `upload_index`. A rejected
* batch carries one `errors` entry per file, so each chip can say why it
* failed instead of every chip repeating the first file's reason.
*/
export function parseUploadErrorsByIndex(body: string): Map<number, string> {
const byIndex = new Map<number, string>();
if (!body) return byIndex;
try {
const parsed = JSON.parse(body) as { errors?: unknown };
if (!Array.isArray(parsed?.errors)) return byIndex;
for (const entry of parsed.errors as {
upload_index?: unknown;
error?: unknown;
}[]) {
if (
typeof entry?.upload_index === 'number' &&
typeof entry.error === 'string'
)
byIndex.set(entry.upload_index, entry.error);
}
} catch {
// Not JSON (a proxy's HTML 502, say) — the caller falls back.
}
return byIndex;
}
+2 -1
View File
@@ -1101,7 +1101,8 @@
"waitingToSend_one": "Will send when {{count}} file finishes processing…",
"waitingToSend_other": "Will send when {{count}} files finish processing…",
"cancelQueuedSend": "Cancel",
"sendBlockedByFailed": "{{names}} couldn't be processed — remove to send"
"sendBlockedByFailed": "{{names}} couldn't be processed — remove to send",
"unsupportedType": "Not a supported file type ({{extension}})"
},
"retry": "Retry",
"reasoning": "Reasoning",
+2
View File
@@ -46,6 +46,8 @@ export interface Attachment {
*/
attachmentId?: string;
token_count?: number;
/** Why a ``failed`` attachment failed, when known (server or client gate). */
errorMessage?: string;
}
export type UploadTaskStatus =
+11 -6
View File
@@ -377,17 +377,18 @@ function Configure-Embeddings {
Write-Host ""
Write-ColorText "Embeddings Configuration" -ForegroundColor "White" -Bold
Write-ColorText "Choose your embeddings provider:" -ForegroundColor "White"
Write-ColorText "1) HuggingFace (default, local)" -ForegroundColor "Yellow"
Write-ColorText "1) Granite multilingual (default, local)" -ForegroundColor "Yellow"
Write-ColorText "2) OpenAI Embeddings" -ForegroundColor "Yellow"
Write-ColorText "3) Custom Remote Embeddings (OpenAI-compatible API)" -ForegroundColor "Yellow"
Write-ColorText "4) all-mpnet-base-v2 (legacy local, English-only)" -ForegroundColor "Yellow"
Write-ColorText "b) Back" -ForegroundColor "Yellow"
Write-Host ""
$emb_choice = Read-Host "Choose option (1-3, or b)"
$emb_choice = Read-Host "Choose option (1-4, or b)"
switch ($emb_choice) {
"1" {
"EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2" | Add-Content -Path $ENV_FILE -Encoding utf8
Write-ColorText "Embeddings set to HuggingFace (local)." -ForegroundColor "Green"
"EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2" | Add-Content -Path $ENV_FILE -Encoding utf8
Write-ColorText "Embeddings set to granite-311m-multilingual-r2 (local)." -ForegroundColor "Green"
}
"2" {
"EMBEDDINGS_NAME=openai_text-embedding-ada-002" | Add-Content -Path $ENV_FILE -Encoding utf8
@@ -404,6 +405,10 @@ function Configure-Embeddings {
if ($emb_key) { "EMBEDDINGS_KEY=$emb_key" | Add-Content -Path $ENV_FILE -Encoding utf8 }
Write-ColorText "Custom remote embeddings configured." -ForegroundColor "Green"
}
"4" {
"EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2" | Add-Content -Path $ENV_FILE -Encoding utf8
Write-ColorText "Embeddings set to all-mpnet-base-v2 (legacy local)." -ForegroundColor "Green"
}
{$_ -eq "b" -or $_ -eq "B"} { return }
default {
Write-Host ""
@@ -742,7 +747,7 @@ function Serve-LocalOllama {
"LLM_NAME=$model_name" | Add-Content -Path $ENV_FILE -Encoding utf8
"VITE_API_STREAMING=true" | Add-Content -Path $ENV_FILE -Encoding utf8
"OPENAI_BASE_URL=http://ollama:11434/v1" | Add-Content -Path $ENV_FILE -Encoding utf8
"EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2" | Add-Content -Path $ENV_FILE -Encoding utf8
"EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2" | Add-Content -Path $ENV_FILE -Encoding utf8
Write-ColorText ".env file configured for Ollama ($($docker_compose_file_suffix.ToUpper()))." -ForegroundColor "Green"
Write-ColorText "Note: MODEL_NAME is set to '$model_name'. You can change it later in the .env file." -ForegroundColor "Yellow"
@@ -907,7 +912,7 @@ function Connect-LocalInferenceEngine {
"LLM_NAME=$model_name" | Add-Content -Path $ENV_FILE -Encoding utf8
"VITE_API_STREAMING=true" | Add-Content -Path $ENV_FILE -Encoding utf8
"OPENAI_BASE_URL=$openai_base_url" | Add-Content -Path $ENV_FILE -Encoding utf8
"EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2" | Add-Content -Path $ENV_FILE -Encoding utf8
"EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2" | Add-Content -Path $ENV_FILE -Encoding utf8
Write-ColorText ".env file configured for $engine_name with OpenAI API format." -ForegroundColor "Green"
Write-ColorText "Note: MODEL_NAME is set to '$model_name'. You can change it later in the .env file." -ForegroundColor "Yellow"
+11 -6
View File
@@ -250,17 +250,18 @@ configure_vector_store() {
configure_embeddings() {
echo -e "\n${DEFAULT_FG}${BOLD}Embeddings Configuration${NC}"
echo -e "${DEFAULT_FG}Choose your embeddings provider:${NC}"
echo -e "${YELLOW}1) HuggingFace (default, local)${NC}"
echo -e "${YELLOW}1) Granite multilingual (default, local)${NC}"
echo -e "${YELLOW}2) OpenAI Embeddings${NC}"
echo -e "${YELLOW}3) Custom Remote Embeddings (OpenAI-compatible API)${NC}"
echo -e "${YELLOW}4) all-mpnet-base-v2 (legacy local, English-only)${NC}"
echo -e "${YELLOW}b) Back${NC}"
echo
read -p "$(echo -e "${DEFAULT_FG}Choose option (1-3, or b): ${NC}")" emb_choice
read -p "$(echo -e "${DEFAULT_FG}Choose option (1-4, or b): ${NC}")" emb_choice
case "$emb_choice" in
1)
echo "EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2" >> "$ENV_FILE"
echo -e "${GREEN}Embeddings set to HuggingFace (local).${NC}"
echo "EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2" >> "$ENV_FILE"
echo -e "${GREEN}Embeddings set to granite-311m-multilingual-r2 (local).${NC}"
;;
2)
echo "EMBEDDINGS_NAME=openai_text-embedding-ada-002" >> "$ENV_FILE"
@@ -277,6 +278,10 @@ configure_embeddings() {
[ -n "$emb_key" ] && echo "EMBEDDINGS_KEY=$emb_key" >> "$ENV_FILE"
echo -e "${GREEN}Custom remote embeddings configured.${NC}"
;;
4)
echo "EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2" >> "$ENV_FILE"
echo -e "${GREEN}Embeddings set to all-mpnet-base-v2 (legacy local).${NC}"
;;
b|B) return ;;
*) echo -e "\n${RED}Invalid choice.${NC}" ; sleep 1 ;;
esac
@@ -541,7 +546,7 @@ serve_local_ollama() {
echo "LLM_NAME=$model_name" >> "$ENV_FILE"
echo "VITE_API_STREAMING=true" >> "$ENV_FILE"
echo "OPENAI_BASE_URL=http://ollama:11434/v1" >> "$ENV_FILE"
echo "EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2" >> "$ENV_FILE"
echo "EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2" >> "$ENV_FILE"
echo -e "${GREEN}.env file configured for Ollama ($(echo "$docker_compose_file_suffix" | tr '[:lower:]' '[:upper:]')${NC}${GREEN}).${NC}"
prompt_advanced_settings
@@ -654,7 +659,7 @@ connect_local_inference_engine() {
echo "LLM_NAME=$model_name" >> "$ENV_FILE"
echo "VITE_API_STREAMING=true" >> "$ENV_FILE"
echo "OPENAI_BASE_URL=$openai_base_url" >> "$ENV_FILE"
echo "EMBEDDINGS_NAME=huggingface_sentence-transformers/all-mpnet-base-v2" >> "$ENV_FILE"
echo "EMBEDDINGS_NAME=ibm-granite/granite-embedding-311m-multilingual-r2" >> "$ENV_FILE"
echo -e "${GREEN}.env file configured for ${BOLD}${engine_name}${NC}${GREEN} with OpenAI API format.${NC}"
echo -e "${YELLOW}Note: MODEL_NAME is set to '${BOLD}$model_name${NC}${YELLOW}'. You can change it later in the .env file.${NC}"
+131
View File
@@ -2369,6 +2369,20 @@ class TestGetStoreAttachmentUserError:
msg = _get_store_attachment_user_error(RuntimeError("oops"))
assert msg == "Failed to process file"
@pytest.mark.unit
def test_unsupported_type_message_is_rebuilt_from_the_filename(self):
"""Nothing is read off the exception, so its text cannot reach a response."""
from application.api.user.attachments.routes import (
_get_store_attachment_user_error,
)
from application.upload_limits import UnsupportedUploadTypeError
err = UnsupportedUploadTypeError("internals: /srv/app/tmp/staged-42")
msg = _get_store_attachment_user_error(err, "clip.mp4")
assert msg == "Unsupported file type: .mp4"
assert "/srv/app" not in msg
class TestRequireLiveSttRedisUnavailable:
"""Cover lines 99-102: Redis unavailable returns 503."""
@@ -2386,3 +2400,120 @@ class TestRequireLiveSttRedisUnavailable:
result = _require_live_stt_redis()
assert hasattr(result, "status_code")
assert result.status_code == 503
MP4_BYTES = b"\x00\x00\x00\x18ftypisom\x00\x00\x02\x00isomiso2avc1mp41\x00\x00\x00\x08free"
class TestStoreAttachmentTypeGate:
"""Unparseable uploads are refused before anything is stored or queued."""
@patch("application.api.user.tasks.store_attachment.delay")
def test_rejects_unsupported_file_type_before_storage(
self, mock_store_attachment, flask_app
):
from application.api.user.attachments.routes import StoreAttachment
app = Flask(__name__)
mock_storage = MagicMock()
with patch("application.api.user.base.storage", mock_storage), app.test_request_context(
"/api/store_attachment",
method="POST",
data={"file": (io.BytesIO(MP4_BYTES), "clip.mp4")},
content_type="multipart/form-data",
):
request.decoded_token = {"sub": "test_user"}
response = StoreAttachment().post()
assert _get_response_status(response) == 400
body = _get_response_json(response)
assert body["success"] is False
assert body["message"] == "Unsupported file type: .mp4"
assert body["errors"] == [
{"upload_index": 0, "filename": "clip.mp4", "error": "Unsupported file type: .mp4"}
]
mock_storage.save_file.assert_not_called()
mock_store_attachment.assert_not_called()
@patch("application.api.user.tasks.store_attachment.delay")
def test_rejects_a_binary_renamed_to_a_text_suffix(
self, mock_store_attachment, flask_app
):
""".txt has no parser, so it is judged on content like any other suffix."""
from application.api.user.attachments.routes import StoreAttachment
app = Flask(__name__)
mock_storage = MagicMock()
with patch("application.api.user.base.storage", mock_storage), app.test_request_context(
"/api/store_attachment",
method="POST",
data={"file": (io.BytesIO(MP4_BYTES), "notes.txt")},
content_type="multipart/form-data",
):
request.decoded_token = {"sub": "test_user"}
response = StoreAttachment().post()
assert _get_response_status(response) == 400
assert _get_response_json(response)["message"] == "Unsupported file type: .txt"
mock_storage.save_file.assert_not_called()
mock_store_attachment.assert_not_called()
@patch("application.api.user.tasks.store_attachment.delay")
def test_accepts_a_text_file_with_no_dedicated_parser(
self, mock_store_attachment, flask_app
):
"""A .py or .log is read by the plain-text fallthrough and must stay allowed."""
from application.api.user.attachments.routes import StoreAttachment
app = Flask(__name__)
mock_storage = MagicMock()
mock_storage.save_file.return_value = {"storage_type": "local"}
mock_store_attachment.return_value = SimpleNamespace(id="task-py")
with patch("application.api.user.base.storage", mock_storage), app.test_request_context(
"/api/store_attachment",
method="POST",
data={"file": (io.BytesIO(b"def main():\n return 1\n"), "main.py")},
content_type="multipart/form-data",
):
request.decoded_token = {"sub": "test_user"}
response = StoreAttachment().post()
assert _get_response_status(response) == 200
assert _get_response_json(response)["task_id"] == "task-py"
assert mock_store_attachment.call_count == 1
@patch("application.api.user.tasks.store_attachment.delay")
def test_batch_skips_unsupported_file_and_keeps_the_rest(
self, mock_store_attachment, flask_app
):
from application.api.user.attachments.routes import StoreAttachment
app = Flask(__name__)
mock_storage = MagicMock()
mock_storage.save_file.return_value = {"storage_type": "local"}
mock_store_attachment.return_value = SimpleNamespace(id="task-notes")
with patch("application.api.user.base.storage", mock_storage), app.test_request_context(
"/api/store_attachment",
method="POST",
data={
"file": [
(io.BytesIO(MP4_BYTES), "clip.mp4"),
(io.BytesIO(b"hello"), "notes.txt"),
]
},
content_type="multipart/form-data",
):
request.decoded_token = {"sub": "test_user"}
response = StoreAttachment().post()
assert _get_response_status(response) == 200
body = _get_response_json(response)
assert [task["upload_index"] for task in body["tasks"]] == [1]
assert body["tasks"][0]["filename"] == "notes.txt"
assert body["errors"] == [
{"upload_index": 0, "filename": "clip.mp4", "error": "Unsupported file type: .mp4"}
]
assert mock_storage.save_file.call_count == 1
assert mock_store_attachment.call_count == 1
@@ -95,6 +95,12 @@ class TestCreateWikiSource:
assert row["type"] == "wiki"
assert row["config"]["kind"] == "wiki"
assert int(row["tokens"]) == 0
# Wiki pages get embedded like any other source, so the row has to name
# the model that did it. NULL reads as "the legacy model" to the boot
# mismatch check, which then reports the source as stale forever.
from application.core.settings import settings
assert row["model"] == settings.EMBEDDINGS_NAME
# No seed content → no re-embed, and never any ingest/reingest task.
mock_reembed.assert_not_called()
mock_ingest.assert_not_called()
+85
View File
@@ -769,3 +769,88 @@ class TestSubmitFeedbackHappy:
response = SubmitFeedback().post()
assert response.status_code == 400
@pytest.mark.unit
class TestSubmitFeedbackWithApiKey:
"""api_key callers carry no JWT."""
def _seed_agent_with_key(self, pg_conn, owner, key):
from application.storage.db.repositories.agents import AgentsRepository
return AgentsRepository(pg_conn).create(owner, "widget", "published", key=key)
def _post(self, app, pg_conn, payload):
from application.api.user.conversations.routes import SubmitFeedback
with _patch_conversations_db(pg_conn), app.test_request_context("/api/feedback", method="POST", json=payload):
from flask import request
request.decoded_token = None
return SubmitFeedback().post()
def test_valid_key_rates_its_own_conversation(self, app, pg_conn):
from application.storage.db.repositories.conversations import (
ConversationsRepository,
)
owner, key = "owner-fb-key", "agent-key-ok"
self._seed_agent_with_key(pg_conn, owner, key)
repo = ConversationsRepository(pg_conn)
conv_id = str(repo.create(owner, name="via widget", api_key=key)["id"])
repo.append_message(conv_id, {"prompt": "p", "response": "r"})
response = self._post(
app,
pg_conn,
{
"feedback": "LIKE",
"question_index": 0,
"conversation_id": conv_id,
"api_key": key,
},
)
assert response.status_code == 200
fb = repo.get_messages(conv_id)[0].get("feedback")
assert fb and fb.get("text") == "like"
def test_key_cannot_rate_owner_conversation_it_did_not_create(self, app, pg_conn):
from application.storage.db.repositories.conversations import (
ConversationsRepository,
)
owner, key = "owner-fb-scope", "agent-key-scope"
self._seed_agent_with_key(pg_conn, owner, key)
# Owner's conversation, not created with the key.
conv_id = _seed_conversation(pg_conn, owner, name="owner private")
ConversationsRepository(pg_conn).append_message(conv_id, {"prompt": "p", "response": "r"})
response = self._post(
app,
pg_conn,
{
"feedback": "LIKE",
"question_index": 0,
"conversation_id": conv_id,
"api_key": key,
},
)
assert response.status_code == 404
def test_unknown_key_is_unauthorized(self, app, pg_conn):
conv_id = _seed_conversation(pg_conn, "owner-fb-bad")
response = self._post(
app,
pg_conn,
{
"feedback": "LIKE",
"question_index": 0,
"conversation_id": conv_id,
"api_key": "no-such-key",
},
)
assert response.status_code == 401
+21 -4
View File
@@ -62,16 +62,18 @@ from sqlalchemy import create_engine
_ALEMBIC_INI = Path(__file__).resolve().parent.parent / "application" / "alembic.ini"
def _migrate_template_db(host, port, user, dbname, password, **_loader_options) -> None:
def _migrate_template_db(host, port, user, dbname, password, **kwargs) -> None:
"""Run alembic ``upgrade head`` into the session's template database.
Called once per session by pytest-postgresql's ``load=`` hook (against
the template DB, before any test runs). Runs in a subprocess so the
parent process never imports application settings with this URI cached.
``**_loader_options`` absorbs keyword arguments newer pytest-postgresql
releases pass to loaders (9.0.0 added ``autocommit=``); alembic manages
its own connection, so they are irrelevant here.
``**kwargs`` swallows keywords newer plugin releases hand to loaders --
9.0.0 added ``autocommit``, which broke every DB test on a plugin bump
alone. They describe the connection the caller already gave us, so
ignoring them keeps one signature working across versions instead of
pinning the plugin back.
"""
url = (
f"postgresql+psycopg://{user}:{password or ''}@{host}:{port}/{dbname}"
@@ -162,6 +164,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)
+21
View File
@@ -514,6 +514,27 @@ class TestEmbeddingDim:
store = GraphStore.__new__(GraphStore)
assert store._embedding_dim() == 1536
def test_none_dimension_falls_back_to_default(self, monkeypatch):
"""A remote model outside the registry reports ``None``, not nothing.
``getattr`` with a default cannot catch that -- the attribute exists --
so the width reached the DDL as ``vector(None)``.
"""
from application.vectorstore import base as base_module
monkeypatch.setattr(base_module.settings, "EMBEDDINGS_BASE_URL", None)
fake_embedding = MagicMock()
fake_embedding.dimension = None
monkeypatch.setattr(
base_module.EmbeddingsSingleton,
"get_instance",
staticmethod(lambda *a, **k: fake_embedding),
)
monkeypatch.setattr(GraphStore, "_embedding_dim", _REAL_EMBEDDING_DIM)
store = GraphStore.__new__(GraphStore)
assert store._embedding_dim() == store_module.DEFAULT_NAME_EMBEDDING_DIM
def test_falls_back_to_default_dimension(self, monkeypatch):
from application.vectorstore import base as base_module
@@ -179,8 +179,12 @@ class TestPerRoundUsageRows:
list(handler.handle_streaming(agent, first, {}, []))
assert len(rows) == 2
assert rows[0] == {"prompt_tokens": 100, "generated_tokens": 10}
assert rows[1] == {"prompt_tokens": 200, "generated_tokens": 20}
# Provider-reported counts replace the estimates; the rest of the
# call record (e.g. ``model``) rides along, hence a subset check
# rather than dict equality.
counts = lambda row: {k: row[k] for k in ("prompt_tokens", "generated_tokens")} # noqa: E731
assert counts(rows[0]) == {"prompt_tokens": 100, "generated_tokens": 10}
assert counts(rows[1]) == {"prompt_tokens": 200, "generated_tokens": 20}
class TestPreferProviderUsageClaim:
+35
View File
@@ -2347,3 +2347,38 @@ class TestBuildConversationFromMessagesEmpty:
result = handler._build_conversation_from_messages(messages)
assert result is not None
assert len(result.get("queries", [])) >= 1
class TestFinishedLogCacheBins:
"""``cached_tokens`` / ``cache_write_tokens`` ride on the finish events
only when a provider reported them, so dashboards can compute a cache
hit rate without a sentinel zero polluting providers that never report."""
def test_stream_finished_carries_cache_bins_when_given(self, caplog):
import logging as _logging
llm = StubLLM()
with caplog.at_level(_logging.INFO, logger="root"):
llm._emit_stream_finished_log(
"m1",
prompt_tokens=100,
completion_tokens=5,
latency_ms=10,
cached_tokens=80,
cache_write_tokens=7,
)
evt = next(r for r in caplog.records if r.message == "llm_stream_finished")
assert evt.cached_tokens == 80
assert evt.cache_write_tokens == 7
def test_gen_finished_omits_cache_bins_when_none(self, caplog):
import logging as _logging
llm = StubLLM()
with caplog.at_level(_logging.INFO, logger="root"):
llm._emit_gen_finished_log(
"m1", prompt_tokens=100, completion_tokens=5, latency_ms=10
)
evt = next(r for r in caplog.records if r.message == "llm_gen_finished")
assert not hasattr(evt, "cached_tokens")
assert not hasattr(evt, "cache_write_tokens")
+56
View File
@@ -1550,3 +1550,59 @@ def test_responses_gen_stream_retries_without_tools(monkeypatch):
assert len(create.calls) == 2
assert "tools" in create.calls[0]
assert "tools" not in create.calls[1]
@pytest.mark.unit
def test_record_chat_usage_captures_cache_write_bin(monkeypatch):
"""Newer OpenAI-family deployments report cache writes next to cache reads."""
llm = _make_llm(monkeypatch, ModelCapabilities(api_flavor="chat_completions"))
llm._record_chat_usage(_ns(
prompt_tokens=10,
completion_tokens=7,
total_tokens=17,
prompt_tokens_details=_ns(cached_tokens=4, cache_write_tokens=6),
completion_tokens_details=None,
))
assert llm._last_usage["prompt_tokens_details"] == {
"cached_tokens": 4,
"cache_write_tokens": 6,
}
@pytest.mark.unit
def test_record_responses_metadata_captures_cache_write_bin(monkeypatch):
llm = _make_llm(monkeypatch, _responses_caps())
llm._record_responses_metadata(_ns(
id="resp_1",
usage=_ns(
input_tokens=10,
output_tokens=7,
total_tokens=17,
input_tokens_details=_ns(cached_tokens=0, cache_write_tokens=9),
output_tokens_details=None,
),
))
# A pure cache-write turn (first request of a prefix) has reads=0 and
# writes>0; both bins were reported, so both are recorded — a reported
# zero is "no cache hits", not "unknown".
assert llm._last_usage["prompt_tokens_details"] == {
"cached_tokens": 0,
"cache_write_tokens": 9,
}
@pytest.mark.unit
def test_record_responses_metadata_omits_unreported_cache_bins(monkeypatch):
"""A provider that says nothing about caching must persist as NULL, not 0."""
llm = _make_llm(monkeypatch, _responses_caps())
llm._record_responses_metadata(_ns(
id="resp_2",
usage=_ns(
input_tokens=10,
output_tokens=7,
total_tokens=17,
input_tokens_details=None,
output_tokens_details=None,
),
))
assert "prompt_tokens_details" not in llm._last_usage
+23
View File
@@ -0,0 +1,23 @@
import pytest
@pytest.fixture(scope="session")
def _mpnet_tokenizer():
"""Real WordPiece tokenizer; skipped when the hub is unreachable."""
pytest.importorskip("tokenizers")
from tokenizers import Tokenizer
try:
tokenizer = Tokenizer.from_pretrained("sentence-transformers/all-mpnet-base-v2")
except Exception as exc: # offline CI
pytest.skip(f"tokenizer unavailable: {exc}")
tokenizer.no_padding()
tokenizer.no_truncation()
return tokenizer
@pytest.fixture
def hf_counter(_mpnet_tokenizer):
from application.parser.tokenization import HuggingFaceCounter
return HuggingFaceCounter(_mpnet_tokenizer, "sentence-transformers/all-mpnet-base-v2")
+109
View File
@@ -0,0 +1,109 @@
"""Tests for the shared file-extension constants."""
import re
from importlib.util import find_spec
from pathlib import Path
import pytest
from application.parser.file.bulk import get_default_file_extractor
from application.parser.file.constants import (
ATTACHMENT_PARSER_EXTENSIONS,
attachment_extension,
has_attachment_parser,
SUPPORTED_SOURCE_EXTENSIONS,
)
FRONTEND_CONSTANTS = (
Path(__file__).resolve().parents[3] / "frontend/src/constants/fileUpload.ts"
)
@pytest.mark.unit
def test_parser_extensions_match_the_extractor():
"""The list is a hand-kept literal; this is what keeps it honest.
``ATTACHMENT_PARSER_EXTENSIONS`` cannot import the extractor (that would
pull docling into the API process), so a parser added to
``get_default_file_extractor`` would otherwise stay refused as an
attachment — which is how .webp, .tiff, .vtt and .xml were first missed.
A suffix listed here but *without* a parser is the opposite failure: it
skips the content check and reaches the plain-text fallthrough, which is
how a binary named notes.txt got through.
"""
parser_suffixes = set(get_default_file_extractor())
assert parser_suffixes <= set(ATTACHMENT_PARSER_EXTENSIONS), (
"parsers exist for suffixes the attachment gate refuses: "
f"{sorted(parser_suffixes - set(ATTACHMENT_PARSER_EXTENSIONS))}"
)
if find_spec("docling") is not None:
assert parser_suffixes == set(ATTACHMENT_PARSER_EXTENSIONS), (
"listed as parser-backed but nothing parses them: "
f"{sorted(set(ATTACHMENT_PARSER_EXTENSIONS) - parser_suffixes)}"
)
@pytest.mark.unit
def test_frontend_mirrors_the_parser_extensions():
"""The composer gates uploads client-side; a divergent list refuses valid files.
``ATTACHMENT_PARSER_EXTENSIONS`` in the frontend is a hand-kept copy of
the backend's. Nothing but this test connects them, and .mdx, .xhtml and
.adoc were blocked in the UI while the API accepted them before it
existed.
"""
if not FRONTEND_CONSTANTS.exists():
pytest.skip("frontend sources not present in this checkout")
source = FRONTEND_CONSTANTS.read_text(encoding="utf-8")
match = re.search(
r"ATTACHMENT_PARSER_EXTENSIONS:\s*readonly string\[\]\s*=\s*\[(.*?)\]",
source,
re.DOTALL,
)
assert match, "ATTACHMENT_PARSER_EXTENSIONS not found in fileUpload.ts"
frontend_extensions = set(re.findall(r"'(\.[^']+)'", match.group(1)))
assert frontend_extensions == set(ATTACHMENT_PARSER_EXTENSIONS), (
"frontend and backend attachment lists disagree — "
f"backend only: {sorted(set(ATTACHMENT_PARSER_EXTENSIONS) - frontend_extensions)}, "
f"frontend only: {sorted(frontend_extensions - set(ATTACHMENT_PARSER_EXTENSIONS))}"
)
@pytest.mark.unit
def test_source_ingestion_types_are_parser_backed_apart_from_text():
"""Source ingestion's own list must not smuggle a parserless suffix past the sniff."""
assert set(SUPPORTED_SOURCE_EXTENSIONS) - set(ATTACHMENT_PARSER_EXTENSIONS) == {
".txt"
}
@pytest.mark.unit
@pytest.mark.parametrize(
("filename", "expected"),
[
("Photo.JPG", ".jpg"),
("/tmp/dir.d/report.PDF", ".pdf"),
("archive.tar.gz", ".gz"),
("Dockerfile", ""),
(".env", ""),
("trailing.", "."),
(None, ""),
],
)
def test_attachment_extension(filename, expected):
assert attachment_extension(filename) == expected
@pytest.mark.unit
def test_has_attachment_parser_is_case_insensitive_and_excludes_zip():
assert has_attachment_parser("Report.PDF")
assert has_attachment_parser("scan.WebP")
# No parser: admitted (or not) on content instead, never on the suffix.
assert not has_attachment_parser("archive.zip")
assert not has_attachment_parser("main.py")
assert not has_attachment_parser("clip.mp4")
# .txt is the plain-text fallthrough itself, not a parser.
assert not has_attachment_parser("notes.txt")
+74 -3
View File
@@ -108,18 +108,89 @@ class TestSplitDocument:
result = chunker.split_document(doc)
assert len(result) > 1
# First chunk should contain header
assert "h1" in result[0].text
# Only the first: this is what duplicate_headers=False means.
assert all("h1" not in chunk.text for chunk in result[1:])
def test_split_duplicates_header(self):
"""Every chunk carries the header when the flag is set.
Asserting only the first chunk passed even while the flag did nothing:
the implementation cleared the header after the first iteration, so
duplicate_headers was unreachable.
"""
chunker = Chunker(max_tokens=50, min_tokens=5, duplicate_headers=True)
text = "h1\nh2\nh3\n" + "word " * 200
doc = Document(text=text, doc_id="doc1")
result = chunker.split_document(doc)
assert len(result) > 1
# First chunk should contain header
assert "h1" in result[0].text
assert all("h1" in chunk.text for chunk in result)
def test_oversized_header_does_not_multiply_chunks(self):
"""A header past the budget used to leave a one-token body budget.
``max(1, max_tokens - header_tokens)`` bottomed out at 1, so the body
was cut into one chunk per token -- each still over ``max_tokens``,
since the whole header was prepended to it.
"""
chunker = Chunker(max_tokens=50, min_tokens=5, duplicate_headers=True)
header = ("verylongheaderword " * 40) + "h2\nh3\n"
body = "word " * 100
doc = Document(text=f"{header}\n{body}", doc_id="doc1")
result = chunker.split_document(doc)
body_tokens = chunker.counter.count(body)
assert len(result) <= body_tokens // 10, "chunk count must track the budget"
for chunk in result:
assert chunk.extra_info["token_count"] <= chunker.max_tokens
assert "".join(chunk.text for chunk in result) == doc.text
def test_header_taking_most_of_the_budget_is_not_duplicated(self):
"""Duplication is dropped when it would leave almost no room for body.
Below a quarter of the budget every chunk is mostly repeated header,
which multiplies the chunk count without adding retrievable text.
"""
chunker = Chunker(max_tokens=50, min_tokens=5, duplicate_headers=True)
header = ("headerword " * 20) + "h2\nh3\n"
doc = Document(text=f"{header}\n" + "word " * 200, doc_id="doc1")
result = chunker.split_document(doc)
assert len(result) > 1
assert "headerword" in result[0].text
assert all("headerword" not in chunk.text for chunk in result[1:])
def test_header_only_document_is_not_dropped(self):
"""The loop only ever emitted the header attached to a body piece.
With no body there was no piece, so an entire document disappeared
from the index with no error and no log line.
"""
chunker = Chunker(max_tokens=50, min_tokens=1, duplicate_headers=False)
text = "h1\nh2\nh3\n"
doc = Document(text=text, doc_id="doc1")
result = chunker.split_document(doc)
assert len(result) == 1
assert result[0].text == text
def test_header_only_document_keeps_its_text_through_chunk(self):
"""The reachable shape: three lines that alone exceed the budget."""
chunker = Chunker(max_tokens=5, min_tokens=1, duplicate_headers=False)
text = "h1\nh2\nh3\n"
result = chunker.chunk([Document(text=text, doc_id="doc1")])
assert result, "the document must not vanish"
assert "".join(chunk.text for chunk in result) == text
def test_zero_max_tokens_does_not_split_per_token(self):
chunker = Chunker(max_tokens=0, min_tokens=1)
assert chunker.max_tokens == 1
def test_split_preserves_embedding(self):
chunker = Chunker(max_tokens=50, min_tokens=5)
+4 -2
View File
@@ -15,11 +15,13 @@ from application.parser.chunking_strategies import (
SemanticChunker,
)
from application.parser.schema.base import Document
from application.utils import get_encoding
from application.parser.tokenization import get_token_counter
def _tok(text: str) -> int:
return len(get_encoding().encode(text))
# Chunk budgets are enforced in the embedding model's tokenizer, so a
# test that measures them in cl100k measures the wrong thing.
return get_token_counter().count(text)
@pytest.mark.unit
+300
View File
@@ -0,0 +1,300 @@
"""Chunk sizes must be counted in the embedding model's units, and splitting
must never rewrite the text it splits."""
import pytest
from application.parser import tokenization
from application.parser.tokenization import (
HuggingFaceCounter,
TiktokenCounter,
get_token_counter,
)
SAMPLES = [
"Hello World: DocsGPT ANSWERS Questions.",
"The quick brown fox jumps over the lazy dog. " * 40,
"Comment configurer l'authentification avec une clé API ?",
"def embed(text: str) -> list[float]:\n return model.encode(text)\n",
"Ünïcödé — em-dashes, curly “quotes”, and 日本語 text.",
"a,b,c\n1,2,3\n4,5,6\n" * 30,
]
class _StubEncoding:
"""Whitespace tokenizer standing in for tiktoken."""
def encode_ordinary(self, text):
return [ord(c) for c in text]
def decode(self, ids):
return "".join(chr(i) for i in ids)
def decode_with_offsets(self, ids):
# One token per character, so each token starts where the last ended.
return self.decode(ids), list(range(len(ids)))
@pytest.fixture(autouse=True)
def _clear_cache():
tokenization.reset_cache()
yield
tokenization.reset_cache()
class TestSplittingPreservesText:
"""The property that protects every stored document."""
@pytest.mark.parametrize("text", SAMPLES)
def test_tiktoken_split_reassembles_exactly(self, text, monkeypatch):
monkeypatch.setattr(tokenization, "get_encoding", _StubEncoding)
counter = TiktokenCounter()
pieces = counter.split(text, 7)
assert "".join(pieces) == text
@pytest.mark.parametrize("text", SAMPLES)
def test_hf_split_reassembles_exactly(self, text, hf_counter):
"""WordPiece lowercases on decode, so splitting must slice, not decode."""
pieces = hf_counter.split(text, 7)
assert "".join(pieces) == text
def test_hf_split_does_not_lowercase(self, hf_counter):
text = "Hello World: DocsGPT ANSWERS Questions."
assert "".join(hf_counter.split(text, 3)) == text
assert "DocsGPT" in "".join(hf_counter.split(text, 3))
@pytest.mark.parametrize("text", SAMPLES)
def test_every_piece_is_within_budget(self, text, hf_counter):
budget = 10
for piece in hf_counter.split(text, budget):
# The final piece can absorb trailing characters the tokenizer
# dropped, so allow a small overshoot there only.
assert hf_counter.count(piece) <= budget + 2
def test_short_text_is_returned_whole(self, hf_counter):
assert hf_counter.split("short", 100) == ["short"]
def test_empty_text_yields_no_pieces(self, hf_counter):
assert hf_counter.split("", 10) == []
def test_zero_budget_is_clamped_not_infinite_loop(self, hf_counter):
pieces = hf_counter.split("some words here to split", 0)
assert "".join(pieces) == "some words here to split"
class TestCounting:
def test_counts_differ_between_tokenizers(self, hf_counter, monkeypatch):
"""The whole point: mpnet and cl100k disagree, so units matter."""
monkeypatch.setattr(tokenization, "get_encoding", _StubEncoding)
text = "internationalisation tokenization"
assert hf_counter.count(text) != TiktokenCounter().count(text)
def test_empty_text_counts_zero(self, hf_counter):
assert hf_counter.count("") == 0
class TestSelection:
def test_registered_model_uses_its_own_tokenizer(self):
counter = get_token_counter("huggingface_sentence-transformers/all-mpnet-base-v2")
assert isinstance(counter, HuggingFaceCounter)
assert counter.name == "sentence-transformers/all-mpnet-base-v2"
def test_openai_model_falls_back_to_cl100k(self):
"""OpenAI models are served remotely and genuinely count cl100k."""
assert isinstance(get_token_counter("openai_text-embedding-ada-002"), TiktokenCounter)
def test_unreachable_tokenizer_falls_back_rather_than_raising(self, monkeypatch):
monkeypatch.setattr(tokenization, "_load_hf_counter", lambda repo: None)
assert isinstance(get_token_counter("granite-311m"), TiktokenCounter)
def test_counter_is_cached_per_model(self):
first = get_token_counter("granite-311m")
assert get_token_counter("granite-311m") is first
class TestTiktokenSplitAgainstRealCl100k:
"""The stub above is one token per character, so it can never place a cut
inside a character. Real cl100k can, and that is the case that corrupted
text: decoding each window on its own turns a straddled multi-byte
character into U+FFFD on both sides of the cut."""
@pytest.fixture
def real_counter(self):
try:
counter = TiktokenCounter()
counter.count("probe")
except Exception as exc: # offline CI, same policy as the HF fixture
pytest.skip(f"cl100k encoding unavailable: {exc}")
return counter
# 2000 is the shipped default max_tokens, 384 mpnet's window; the small
# values place many more cuts per unit of text.
@pytest.mark.parametrize("window", [1, 2, 3, 7, 128, 384, 2000])
@pytest.mark.parametrize(
"text",
[
"日本語のテキストです。絵文字も🎉あります。",
"検索は自然言語でできます。" * 200,
"Здравствуйте, как настроить аутентификацию?",
"🎉🎊✨🚀🔥💡📚🧠" * 50,
"Ünïcödé — em-dashes, curly “quotes”, and 日本語 text.",
],
ids=["ja-short", "ja-long", "ru", "emoji", "mixed"],
)
def test_split_reassembles_exactly(self, real_counter, text, window):
pieces = real_counter.split(text, window)
assert "".join(pieces) == text
@pytest.mark.parametrize("window", [1, 3, 128, 2000])
def test_split_never_emits_a_replacement_character(self, real_counter, window):
text = "検索は自然言語でできます。絵文字も🎉あります。" * 100
assert "�" not in "".join(real_counter.split(text, window))
def test_first_window_budget_is_honoured_and_lossless(self, real_counter):
text = "日本語のテキストです。" * 50
pieces = real_counter.split(text, 20, first_max_tokens=5)
assert "".join(pieces) == text
assert real_counter.count(pieces[0]) <= 5
class TestTiktokenCounterEdges:
"""The cl100k path is the fallback, so its edges matter as much."""
@pytest.fixture
def counter(self, monkeypatch):
monkeypatch.setattr(tokenization, "get_encoding", _StubEncoding)
return TiktokenCounter()
def test_empty_text_counts_zero_and_splits_to_nothing(self, counter):
assert counter.count("") == 0
assert counter.split("", 10) == []
def test_text_within_budget_is_returned_whole(self, counter):
assert counter.split("abc", 10) == ["abc"]
def test_first_window_can_be_smaller_than_the_rest(self, counter):
"""A header eats into the first chunk's budget only."""
pieces = counter.split("abcdefghij", 4, first_max_tokens=2)
assert pieces[0] == "ab"
assert "".join(pieces) == "abcdefghij"
class TestCounterContract:
def test_base_class_requires_an_implementation(self):
base = tokenization.TokenCounter()
with pytest.raises(NotImplementedError):
base.count("x")
with pytest.raises(NotImplementedError):
base.split("x", 1)
class TestFallbackWhenTokenizerUnavailable:
def test_load_failure_returns_none_rather_than_raising(self, monkeypatch, caplog):
"""Chunking must survive an offline host or a bad repo name."""
import builtins
real_import = builtins.__import__
def boom(name, *args, **kwargs):
if name == "tokenizers":
raise ImportError("no tokenizers here")
return real_import(name, *args, **kwargs)
monkeypatch.setattr(builtins, "__import__", boom)
assert tokenization._load_hf_counter("some/repo") is None
def test_selection_falls_back_to_cl100k_on_failure(self, monkeypatch):
monkeypatch.setattr(tokenization, "_load_hf_counter", lambda repo: None)
assert isinstance(get_token_counter("granite-97m"), TiktokenCounter)
def test_reset_cache_forces_reselection(self, monkeypatch):
first = get_token_counter("granite-311m")
tokenization.reset_cache()
monkeypatch.setattr(tokenization, "_load_hf_counter", lambda repo: None)
assert get_token_counter("granite-311m") is not first
class TestOffsetsWithoutSpans:
"""Some tokenizers emit ``(0, 0)`` for specials or normalised-away chars.
Those tokens consume budget but point at no text, so the splitter has to
skip the window rather than emit an empty piece or lose the tail.
"""
class _Encoded:
def __init__(self, offsets):
self.offsets = offsets
self.ids = list(range(len(offsets)))
class _Tokenizer:
def __init__(self, offsets):
self._offsets = offsets
def encode(self, text, add_special_tokens=False):
return TestOffsetsWithoutSpans._Encoded(self._offsets)
def _counter(self, offsets):
return HuggingFaceCounter(self._Tokenizer(offsets), "stub")
def test_span_less_windows_are_skipped_not_emitted_empty(self):
# Two real tokens, then a window of pure (0, 0) padding-like entries.
counter = self._counter([(0, 2), (2, 4), (0, 0), (0, 0)])
pieces = counter.split("abcd", 2)
assert "" not in pieces
assert "".join(pieces) == "abcd"
def test_trailing_text_is_never_dropped(self):
"""Offsets that stop short of the string must not lose the remainder."""
counter = self._counter([(0, 1), (1, 2), (2, 3)])
pieces = counter.split("abcdef", 2)
assert "".join(pieces) == "abcdef"
def test_all_span_less_offsets_still_return_the_text(self):
counter = self._counter([(0, 0), (0, 0), (0, 0)])
assert "".join(counter.split("abc", 1)) == "abc"
class TestUnknownTokenCollapse:
"""A tokenizer that folds a long unbroken run into one ``[UNK]``.
WordPiece gives up on any word longer than ``max_input_chars_per_word``
and emits a single unknown token for it. Counting that as one token makes
a base64 blob or a minified bundle look tiny, so the chunker never splits
it and an oversized chunk reaches the embedding server.
"""
class _CollapsingEncoding:
"""One token per whitespace-separated word, however long the word."""
def __init__(self, text):
self.ids = []
self.offsets = []
cursor = 0
for word in text.split(" "):
if word:
self.ids.append(0)
self.offsets.append((cursor, cursor + len(word)))
cursor += len(word) + 1
class _CollapsingTokenizer:
def encode(self, text, add_special_tokens=False):
return TestUnknownTokenCollapse._CollapsingEncoding(text)
def _counter(self):
return tokenization.HuggingFaceCounter(self._CollapsingTokenizer(), "stub")
def test_long_unbroken_run_is_charged_by_its_span(self):
counter = self._counter()
assert counter.count("a" * 9000) > 100
def test_ordinary_prose_is_unaffected(self):
counter = self._counter()
text = "the quick brown fox jumps over the lazy dog"
assert counter.count(text) == 9
def test_split_bounds_a_collapsed_run(self):
counter = self._counter()
text = "a" * 9000
pieces = counter.split(text, 20)
assert "".join(pieces) == text, "split must not lose or alter text"
assert len(pieces) > 1, "a collapsed run must still be cut into pieces"
assert all(counter.count(p) <= 20 for p in pieces)
+96
View File
@@ -0,0 +1,96 @@
"""Model pre-fetching, which the Docker build runs to bake artifacts in."""
import sys
import types
from unittest.mock import MagicMock, patch
import pytest
from application.scripts import prefetch_models
from application.vectorstore.model_registry import GRANITE_311M, MPNET, OPENAI_ADA_002
@pytest.fixture
def fake_fastembed():
"""Stand in for FastEmbed so nothing is downloaded."""
text_embedding = MagicMock()
pooling = types.SimpleNamespace(CLS="CLS", MEAN="MEAN")
module = types.ModuleType("fastembed")
module.TextEmbedding = text_embedding
desc = types.ModuleType("fastembed.common.model_description")
desc.PoolingType = pooling
desc.ModelSource = lambda hf=None: {"hf": hf}
with patch.dict(
sys.modules,
{"fastembed": module, "fastembed.common.model_description": desc},
):
yield text_embedding
class TestPrefetch:
def test_defaults_cover_legacy_and_new_install(self):
"""An upgraded image must still serve mpnet; a new one needs granite."""
assert MPNET.name in prefetch_models.DEFAULT_MODELS
assert GRANITE_311M.name in prefetch_models.DEFAULT_MODELS
def test_fetches_named_models(self, fake_fastembed):
fetched = prefetch_models.prefetch([GRANITE_311M.name])
assert fetched == [GRANITE_311M.repo]
assert fake_fastembed.call_args.kwargs["model_name"] == GRANITE_311M.repo
def test_registers_with_registry_pooling_and_dim(self, fake_fastembed):
prefetch_models.prefetch([GRANITE_311M.name])
kwargs = fake_fastembed.add_custom_model.call_args.kwargs
assert kwargs["model"] == GRANITE_311M.repo
assert kwargs["pooling"] == "CLS"
assert kwargs["dim"] == GRANITE_311M.dimension
assert kwargs["model_file"] == GRANITE_311M.onnx_file
def test_mean_pooled_model_registers_as_mean(self, fake_fastembed):
prefetch_models.prefetch([MPNET.name])
assert fake_fastembed.add_custom_model.call_args.kwargs["pooling"] == "MEAN"
def test_aliases_resolve(self, fake_fastembed):
assert prefetch_models.prefetch(["granite-311m"]) == [GRANITE_311M.repo]
def test_cache_dir_is_forwarded(self, fake_fastembed):
prefetch_models.prefetch([MPNET.name], cache_dir="/models")
assert fake_fastembed.call_args.kwargs["cache_dir"] == "/models"
def test_cache_dir_omitted_when_absent(self, fake_fastembed):
prefetch_models.prefetch([MPNET.name])
assert "cache_dir" not in fake_fastembed.call_args.kwargs
def test_remote_only_model_is_skipped(self, fake_fastembed):
"""OpenAI embeddings have no local artifacts to cache."""
assert prefetch_models.prefetch([OPENAI_ADA_002.name]) == []
fake_fastembed.assert_not_called()
def test_unknown_model_fails_loudly(self, fake_fastembed):
"""A silent skip at build time is a download at run time, offline."""
with pytest.raises(SystemExit) as excinfo:
prefetch_models.prefetch(["nope/nope"])
assert "nope/nope" in str(excinfo.value)
assert MPNET.name in str(excinfo.value)
def test_several_models_in_one_run(self, fake_fastembed):
fetched = prefetch_models.prefetch([MPNET.name, GRANITE_311M.name])
assert fetched == [MPNET.repo, GRANITE_311M.repo]
class TestMain:
def test_no_args_fetches_the_defaults(self, fake_fastembed):
with patch.object(prefetch_models, "prefetch", return_value=[]) as spy:
assert prefetch_models.main([]) == 0
assert spy.call_args.args[0] == list(prefetch_models.DEFAULT_MODELS)
def test_explicit_args_override_the_defaults(self, fake_fastembed):
with patch.object(prefetch_models, "prefetch", return_value=[]) as spy:
prefetch_models.main(["granite-97m"])
assert spy.call_args.args[0] == ["granite-97m"]
def test_cache_dir_read_from_environment(self, fake_fastembed, monkeypatch):
monkeypatch.setenv("EMBEDDINGS_CACHE_DIR", "/app/models")
with patch.object(prefetch_models, "prefetch", return_value=[]) as spy:
prefetch_models.main(["granite-97m"])
assert spy.call_args.args[1] == "/app/models"
+534
View File
@@ -0,0 +1,534 @@
"""Re-embed script: CLI contract, orchestration, and the FAISS rebuild path."""
from unittest.mock import MagicMock, patch
import pytest
from application.scripts import reembed
def paginating_cursor(chunk_rows, *, graph_rows=(), graph_table=("graph_nodes",)):
"""A cursor answering the reads ``reembed_pgvector`` issues, from memory.
The chunk read is paginated by keyset, so a cursor that returns the same
page for every ``fetchall`` never terminates. Returns the cursor and the
mutable table backing it.
Args:
chunk_rows: ``(id, text)`` rows for the source's chunk table.
graph_rows: ``(id, name)`` rows for ``graph_nodes``.
graph_table: What ``to_regclass`` reports; ``(None,)`` for absent.
Returns:
``(cursor, tables)``, where ``tables["chunks"]`` and
``tables["graph"]`` can be reassigned to change what is served.
"""
tables = {"chunks": list(chunk_rows), "graph": list(graph_rows)}
pending = {"rows": []}
cursor = MagicMock()
def execute(query, params=None):
text = str(query)
if "to_regclass" in text:
pending["rows"] = [graph_table]
elif "FROM graph_nodes" in text:
pending["rows"] = list(tables["graph"])
elif "count(*)" in text:
pending["rows"] = [(len(tables["chunks"]),)]
elif "id > %s" in text:
_, after_id, limit = params
pending["rows"] = [r for r in tables["chunks"] if r[0] > after_id][:limit]
else:
pending["rows"] = list(tables["chunks"])[: params[1]]
cursor.execute.side_effect = execute
cursor.fetchall.side_effect = lambda: pending["rows"]
cursor.fetchone.side_effect = lambda: pending["rows"][0] if pending["rows"] else None
return cursor, tables
class TestCLI:
def test_defaults(self):
args = reembed.build_parser().parse_args([])
assert args.dry_run is False
assert args.sources is None
assert args.batch_size == reembed.DEFAULT_BATCH_SIZE
def test_unsupported_store_exits_without_touching_anything(self):
with patch.object(reembed.settings, "VECTOR_STORE", "qdrant", create=True):
with patch.object(reembed, "run") as run:
assert reembed.main([]) == 2
run.assert_not_called()
def test_supported_stores_are_pgvector_and_faiss(self):
assert set(reembed.SUPPORTED_STORES) == {"pgvector", "faiss"}
def test_sources_are_split_and_trimmed(self):
with patch.object(reembed.settings, "VECTOR_STORE", "faiss", create=True):
with patch.object(reembed, "run", return_value=0) as run:
reembed.main(["--sources", " a , b ,, c "])
assert run.call_args.args[1] == ["a", "b", "c"]
def test_batch_size_is_clamped_to_at_least_one(self):
with patch.object(reembed.settings, "VECTOR_STORE", "faiss", create=True):
with patch.object(reembed, "run", return_value=0) as run:
reembed.main(["--batch-size", "0"])
assert run.call_args.args[2] == 1
class TestRun:
def test_no_sources_is_a_clean_exit(self):
with patch.object(reembed, "list_source_ids", return_value=[]):
assert reembed.run("faiss", None, 8, False) == 0
def test_processes_every_discovered_source(self):
with patch.object(reembed, "list_source_ids", return_value=["a", "b"]):
with patch.object(reembed, "reembed_faiss", return_value=(3, 3)) as handler:
assert reembed.run("faiss", None, 8, False) == 0
assert [call.args[0] for call in handler.call_args_list] == ["a", "b"]
def test_explicit_sources_skip_discovery(self):
with patch.object(reembed, "list_source_ids") as discover:
with patch.object(reembed, "reembed_faiss", return_value=(1, 1)):
reembed.run("faiss", ["only-this"], 8, False)
discover.assert_not_called()
def test_one_failing_source_does_not_stop_the_others(self):
def handler(source_id, batch_size, dry_run):
if source_id == "bad":
raise RuntimeError("boom")
return (1, 1)
with patch.object(reembed, "reembed_faiss", side_effect=handler) as spy:
code = reembed.run("faiss", ["good", "bad", "also-good"], 8, False)
assert code == 1, "a failure must be reported in the exit code"
assert spy.call_count == 3, "later sources must still be attempted"
def test_dry_run_is_reported_as_success(self):
with patch.object(reembed, "reembed_faiss", return_value=(5, 0)):
assert reembed.run("faiss", ["a"], 8, True) == 0
def test_pgvector_uses_the_pgvector_handler(self):
with patch.object(reembed, "reembed_pgvector", return_value=(1, 1)) as handler:
reembed.run("pgvector", ["a"], 8, False)
handler.assert_called_once()
class TestFaissRebuild:
@pytest.fixture
def stores(self):
"""An existing store to read from and the rebuilt one written back."""
existing = MagicMock()
existing.get_chunks.return_value = [
{"doc_id": "1", "text": "alpha", "metadata": {"i": 0}},
{"doc_id": "2", "text": "beta", "metadata": {"i": 1}},
]
rebuilt = MagicMock()
with patch.object(
reembed.VectorCreator, "create_vectorstore", side_effect=[existing, rebuilt]
) as factory:
yield existing, rebuilt, factory
def test_rebuilds_from_stored_text_and_saves(self, stores):
existing, rebuilt, factory = stores
seen, written = reembed.reembed_faiss("s1", batch_size=8, dry_run=False)
assert (seen, written) == (2, 2)
rebuilt.save_local.assert_called_once()
docs = factory.call_args_list[1].kwargs["docs_init"]
assert [d.page_content for d in docs] == ["alpha", "beta"]
assert [d.metadata for d in docs] == [{"i": 0}, {"i": 1}]
def test_embeddings_key_comes_from_settings(self, stores):
"""A placeholder here is sent as the server's bearer token."""
_, _, factory = stores
with patch.object(reembed.settings, "EMBEDDINGS_KEY", "sk-real", create=True):
reembed.reembed_faiss("s1", batch_size=8, dry_run=False)
keys = [call.kwargs.get("embeddings_key") for call in factory.call_args_list]
assert keys == ["sk-real", "sk-real"]
def test_chunk_ids_are_preserved(self, stores):
"""Re-embedding must not renumber chunks.
Fresh ids orphan every GraphRAG ``graph_node_chunks`` row for the
source and invalidate any id a client already holds.
"""
_, _, factory = stores
reembed.reembed_faiss("s1", batch_size=8, dry_run=False)
assert factory.call_args_list[1].kwargs["ids"] == ["1", "2"]
def test_existing_index_is_opened_past_the_dimension_check(self, stores):
"""A width change is the main reason to run this script.
``assert_embedding_dimensions`` refuses to open an index whose width
differs from the configured model -- and its error message recommends
this script, so without the opt-out the advice failed on every source.
The chunk text lives in the sidecar, so reading it needs no match.
"""
_, _, factory = stores
reembed.reembed_faiss("s1", batch_size=8, dry_run=False)
assert factory.call_args_list[0].kwargs["skip_dimension_check"] is True
# The rebuild writes the new width, so it must still be checked.
assert "skip_dimension_check" not in factory.call_args_list[1].kwargs
def test_batch_size_is_forwarded_to_the_rebuild(self, stores):
"""On a remote embeddings server the whole index is otherwise one POST."""
_, _, factory = stores
reembed.reembed_faiss("s1", batch_size=8, dry_run=False)
assert factory.call_args_list[1].kwargs["batch_size"] == 8
def test_existing_index_is_not_deleted(self, stores):
"""The rebuild must not destroy the old index before the new one exists."""
existing, rebuilt, _ = stores
reembed.reembed_faiss("s1", batch_size=8, dry_run=False)
existing.delete_index.assert_not_called()
def test_dry_run_reads_but_never_rebuilds(self):
existing = MagicMock()
existing.get_chunks.return_value = [{"doc_id": "1", "text": "a", "metadata": {}}]
with patch.object(
reembed.VectorCreator, "create_vectorstore", return_value=existing
) as factory:
seen, written = reembed.reembed_faiss("s1", batch_size=8, dry_run=True)
assert (seen, written) == (1, 0)
assert factory.call_count == 1, "no rebuild store may be constructed"
def test_empty_index_is_a_no_op(self):
existing = MagicMock()
existing.get_chunks.return_value = []
with patch.object(
reembed.VectorCreator, "create_vectorstore", return_value=existing
):
assert reembed.reembed_faiss("s1", batch_size=8, dry_run=False) == (0, 0)
class TestPgvectorWithoutTheExtension:
"""Mocked pgvector paths.
The live tests in ``test_reembed_pgvector_live`` skip wherever the cluster
has no pgvector build -- which includes CI -- so the SQL shape and the
batching contract are pinned here too.
"""
@pytest.fixture
def store(self):
"""A store whose cursor serves the paginated reads, not a fixed page.
``reembed_pgvector`` walks the source by keyset, so a cursor that
returns the same rows for every ``fetchall`` never terminates. The fake
answers each of the three queries the function issues from one in-memory
table, which tests mutate through ``rows``.
"""
cursor, table = paginating_cursor([(1, "alpha"), (2, "beta"), (3, "gamma")])
store = MagicMock()
store._table_name = "documents"
store._vector_column = "embedding"
conn = MagicMock()
conn.cursor.return_value = cursor
store._get_connection.return_value = conn
store._embedding.embed_documents.side_effect = lambda texts: [
[0.5] * 4 for _ in texts
]
with patch.object(
reembed.VectorCreator, "create_vectorstore", return_value=store
), patch.object(reembed.settings, "GRAPHRAG_ENABLED", False):
yield store, conn, cursor, table
def test_reads_and_rewrites_every_chunk(self, store):
_, conn, cursor, _ = store
seen, written = reembed.reembed_pgvector("s1", batch_size=64, dry_run=False)
assert (seen, written) == (3, 3)
cursor.executemany.assert_called_once()
conn.commit.assert_called()
def test_pages_are_bounded_by_batch_size(self, store):
"""The read is paginated, so a huge source never lands in memory at once."""
_, _, cursor, table = store
table["chunks"] = [(i, f"chunk-{i}") for i in range(1, 11)]
seen, written = reembed.reembed_pgvector("s1", batch_size=3, dry_run=False)
assert (seen, written) == (10, 10)
selects = [
call for call in cursor.execute.call_args_list
if "SELECT id, text" in str(call.args[0])
]
# Four pages of at most 3, then one empty page to end the walk.
assert len(selects) == 5
assert all(call.args[1][-1] == 3 for call in selects)
def test_dry_run_counts_without_reading_the_text(self, store):
fake_store, _, cursor, _ = store
seen, written = reembed.reembed_pgvector("s1", batch_size=64, dry_run=True)
assert (seen, written) == (3, 0)
fake_store._embedding.embed_documents.assert_not_called()
cursor.executemany.assert_not_called()
assert any(
"count(*)" in str(call.args[0]) for call in cursor.execute.call_args_list
)
def test_batches_commit_separately(self, store):
_, conn, cursor, _ = store
reembed.reembed_pgvector("s1", batch_size=2, dry_run=False)
# 3 rows at batch 2 is two write transactions.
assert cursor.executemany.call_count == 2
assert conn.commit.call_count == 2
def test_failed_batch_rolls_back_and_raises(self, store):
_, conn, cursor, _ = store
cursor.executemany.side_effect = RuntimeError("write failed")
with pytest.raises(RuntimeError):
reembed.reembed_pgvector("s1", batch_size=64, dry_run=False)
conn.rollback.assert_called_once()
def test_connection_is_returned_even_on_failure(self, store):
fake_store, _, cursor, _ = store
cursor.executemany.side_effect = RuntimeError("boom")
with pytest.raises(RuntimeError):
reembed.reembed_pgvector("s1", batch_size=64, dry_run=False)
fake_store.close.assert_called_once()
def test_null_text_does_not_crash_the_embed_call(self, store):
fake_store, _, _, table = store
table["chunks"] = [(1, None), (2, "beta")]
seen, written = reembed.reembed_pgvector("s1", batch_size=64, dry_run=False)
assert (seen, written) == (2, 2)
assert fake_store._embedding.embed_documents.call_args.args[0] == ["", "beta"]
def test_empty_source_is_a_no_op(self, store):
fake_store, _, _, table = store
table["chunks"] = []
assert reembed.reembed_pgvector("s1", batch_size=64, dry_run=False) == (0, 0)
fake_store._embedding.embed_documents.assert_not_called()
def test_source_discovery_returns_sorted_ids(self):
store = MagicMock()
store._table_name = "documents"
cursor = MagicMock()
cursor.fetchall.return_value = [("b",), ("a",)]
conn = MagicMock()
conn.cursor.return_value = cursor
store._get_connection.return_value = conn
with patch.object(
reembed.VectorCreator, "create_vectorstore", return_value=store
):
assert reembed.list_source_ids("pgvector") == ["b", "a"]
class TestFaissSourceDiscovery:
def test_source_ids_come_from_index_directories(self):
storage = MagicMock()
storage.list_files.return_value = [
"indexes/src-a/index.faiss",
"indexes/src-a/index.pkl",
"indexes/src-b/index.faiss",
]
with patch(
"application.storage.storage_creator.StorageCreator.get_storage",
return_value=storage,
):
assert reembed.list_source_ids("faiss") == ["src-a", "src-b"]
def test_storage_failure_is_reported_as_a_usable_error(self):
storage = MagicMock()
storage.list_files.side_effect = OSError("permission denied")
with patch(
"application.storage.storage_creator.StorageCreator.get_storage",
return_value=storage,
):
with pytest.raises(reembed.ReembedError, match="permission denied"):
reembed.list_source_ids("faiss")
def test_unsupported_store_error_names_the_alternatives(self):
with patch.object(reembed.settings, "VECTOR_STORE", "milvus", create=True):
assert reembed.main([]) == 2
class TestGraphNodeReembedding:
"""``graph_nodes.name_embedding`` seeds every graph traversal.
It is written once at extraction time and never revisited, so rewriting
only the chunk table leaves the graph seeding from the previous model --
and since mpnet and granite-311m are both 768-dimensional, the column
accepts the mismatch and nothing reports it.
"""
@pytest.fixture
def graph(self):
store = MagicMock()
store._embedding.embed_documents.side_effect = lambda texts: [
[0.5] * 4 for _ in texts
]
cursor = MagicMock()
cursor.fetchone.return_value = ("graph_nodes",)
cursor.fetchall.return_value = [
("n1", "Alpha"),
("n2", "Beta"),
("n3", "Gamma"),
]
conn = MagicMock()
conn.cursor.return_value = cursor
return store, conn, cursor
def test_rewrites_every_node_name(self, graph):
store, conn, cursor = graph
written = reembed.reembed_graph_nodes(
store, conn, "s1", batch_size=64, dry_run=False
)
assert written == 3
store._embedding.embed_documents.assert_called_once_with(
["Alpha", "Beta", "Gamma"]
)
statement = cursor.executemany.call_args.args[0]
assert "graph_nodes" in statement and "name_embedding" in statement
conn.commit.assert_called()
def test_dry_run_counts_without_embedding(self, graph):
store, conn, cursor = graph
assert reembed.reembed_graph_nodes(store, conn, "s1", 64, dry_run=True) == 3
store._embedding.embed_documents.assert_not_called()
cursor.executemany.assert_not_called()
def test_missing_table_is_a_no_op(self, graph):
store, conn, cursor = graph
cursor.fetchone.return_value = (None,)
assert reembed.reembed_graph_nodes(store, conn, "s1", 64, dry_run=False) == 0
store._embedding.embed_documents.assert_not_called()
def test_batches_commit_separately(self, graph):
store, conn, cursor = graph
reembed.reembed_graph_nodes(store, conn, "s1", batch_size=2, dry_run=False)
assert cursor.executemany.call_count == 2
assert conn.commit.call_count == 2
def test_failed_batch_rolls_back_and_raises(self, graph):
store, conn, cursor = graph
cursor.executemany.side_effect = RuntimeError("write failed")
with pytest.raises(RuntimeError, match="write failed"):
reembed.reembed_graph_nodes(store, conn, "s1", 64, dry_run=False)
conn.rollback.assert_called_once()
def test_pgvector_run_skips_the_graph_when_disabled(self):
store, conn, cursor = self._pgvector_mocks()
with patch.object(
reembed.VectorCreator, "create_vectorstore", return_value=store
), patch.object(reembed.settings, "GRAPHRAG_ENABLED", False):
reembed.reembed_pgvector("s1", batch_size=64, dry_run=False)
assert not self._graph_statements(cursor)
def test_pgvector_run_reembeds_the_graph_when_enabled(self):
store, conn, cursor = self._pgvector_mocks()
with patch.object(
reembed.VectorCreator, "create_vectorstore", return_value=store
), patch.object(reembed.settings, "GRAPHRAG_ENABLED", True):
reembed.reembed_pgvector("s1", batch_size=64, dry_run=False)
assert self._graph_statements(cursor)
@staticmethod
def _pgvector_mocks():
cursor, _ = paginating_cursor(
[(1, "alpha"), (2, "beta")],
graph_rows=[("n1", "Alpha"), ("n2", "Beta")],
)
store = MagicMock()
store._table_name = "documents"
store._vector_column = "embedding"
conn = MagicMock()
conn.cursor.return_value = cursor
store._get_connection.return_value = conn
store._embedding.embed_documents.side_effect = lambda texts: [
[0.5] * 4 for _ in texts
]
return store, conn, cursor
@staticmethod
def _graph_statements(cursor):
return [
call
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
class TestRecordsTheModel:
"""``sources.model`` is what the boot mismatch check reads."""
def test_a_re_embedded_source_is_stamped(self, monkeypatch):
from application.core.settings import settings
monkeypatch.setattr(settings, "EMBEDDINGS_NAME", "new/model", raising=False)
conn = MagicMock()
session = MagicMock()
session.__enter__ = MagicMock(return_value=conn)
session.__exit__ = MagicMock(return_value=False)
with patch("application.storage.db.session.db_session", return_value=session):
reembed.record_source_model("src-1")
params = conn.execute.call_args.args[1]
assert params == {"model": "new/model", "id": "src-1"}
def test_a_dry_run_stamps_nothing(self, monkeypatch):
with patch.object(reembed, "record_source_model") as record, patch.object(
reembed, "list_source_ids", return_value=["a"]
), patch.object(reembed, "reembed_pgvector", return_value=(3, 0)):
reembed.run("pgvector", None, 64, True)
record.assert_not_called()
def test_a_real_run_stamps_each_source(self, monkeypatch):
with patch.object(reembed, "record_source_model") as record, patch.object(
reembed, "list_source_ids", return_value=["a", "b"]
), patch.object(reembed, "reembed_pgvector", return_value=(3, 3)):
reembed.run("pgvector", None, 64, False)
assert [c.args[0] for c in record.call_args_list] == ["a", "b"]
def test_a_failed_source_is_not_stamped(self):
with patch.object(reembed, "record_source_model") as record, patch.object(
reembed, "list_source_ids", return_value=["a"]
), 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()
+297
View File
@@ -0,0 +1,297 @@
"""Live pgvector run of the re-embed script.
Uses the ephemeral pytest-postgresql cluster with a stub embeddings model, so
a real ``UPDATE ... ::vector`` round trip is exercised without downloading a
model. Skips when the cluster has no pgvector build.
"""
from __future__ import annotations
from unittest.mock import patch
import pytest
from application.scripts import reembed
from application.vectorstore import pgvector as pgvector_module
from application.vectorstore.pgvector import PGVectorStore
pytestmark = pytest.mark.integration
DIM = 8
class _Embeddings:
"""Returns a distinct constant per generation, so a rewrite is visible."""
dimension = DIM
def __init__(self, seed: float):
self.seed = seed
self.calls = 0
def embed_documents(self, texts):
self.calls += 1
return [[self.seed] + [0.0] * (DIM - 1) for _ in texts]
def embed_query(self, query):
return [self.seed] + [0.0] * (DIM - 1)
def _dsn(info) -> str:
password = f":{info.password}" if info.password else ""
return f"postgresql://{info.user}{password}@{info.host}:{info.port}/{info.dbname}"
@pytest.fixture(autouse=True)
def _close_pools():
"""Never leak a pool into another test; the DSN dies with the test DB."""
yield
for dsn, pool in list(pgvector_module._POOLS.items()):
try:
pool.close()
except Exception:
# Teardown only: the ephemeral cluster may already be gone, and a
# failure to close a pool for a dead DSN must not fail the test
# that just passed. Dropping the entry below is what matters.
pass
pgvector_module._POOLS.pop(dsn, None)
@pytest.fixture
def live_dsn(postgresql, monkeypatch):
try:
with postgresql.cursor() as cursor:
cursor.execute("CREATE EXTENSION vector;")
postgresql.rollback()
except Exception as exc:
postgresql.rollback()
pytest.skip(f"pgvector extension unavailable: {exc}")
dsn = _dsn(postgresql.info)
from application.core import settings as settings_module
settings = settings_module.settings
monkeypatch.setattr(settings, "VECTOR_STORE", "pgvector", raising=False)
monkeypatch.setattr(settings, "PGVECTOR_CONNECTION_STRING", dsn, raising=False)
monkeypatch.setattr(settings, "GRAPHRAG_ENABLED", False, raising=False)
monkeypatch.setattr(settings, "PGVECTOR_IVFFLAT_PROBES", None, raising=False)
monkeypatch.setattr(settings, "PGVECTOR_POOL_MAX_SIZE", 4, raising=False)
monkeypatch.setattr(settings, "EMBEDDINGS_NAME", "granite-311m", raising=False)
return dsn
def _seed(dsn, source_id, texts, embeddings):
"""Create the schema and insert ``texts`` embedded by ``embeddings``."""
with patch(
"application.vectorstore.base.BaseVectorStore._get_embeddings",
return_value=embeddings,
):
store = PGVectorStore(source_id=source_id, connection_string=dsn)
conn = store._get_connection()
PGVectorStore.create_schema(conn, dimension=DIM)
conn.commit()
store.add_texts(list(texts), metadatas=[{"i": i} for i in range(len(texts))])
store.close()
def _as_list(vector):
"""Normalise a stored vector to a list of floats.
Depending on whether pgvector's adapter is registered on the reading
connection, the value comes back as a ``Vector`` or as its text form
``'[1,0,...]'``. Both are valid; the assertions should not care.
"""
if vector is None:
return []
if hasattr(vector, "to_list"):
return list(vector.to_list())
if isinstance(vector, str):
return [float(part) for part in vector.strip("[]").split(",") if part]
return list(vector)
def _vectors(dsn, source_id):
store = PGVectorStore(source_id=source_id, connection_string=dsn)
conn = store._get_connection()
cursor = conn.cursor()
try:
cursor.execute(
"SELECT text, embedding FROM documents WHERE source_id = %s ORDER BY id",
(source_id,),
)
# pgvector hands back a ``Vector``; normalise to a plain list so the
# assertions read the same whichever adapter is registered.
return [(text, _as_list(vector)) for text, vector in cursor.fetchall()]
finally:
cursor.close()
store.close()
TEXTS = ["alpha document", "beta document", "gamma document"]
class TestReembedPgvectorLive:
def test_rewrites_vectors_and_preserves_text(self, live_dsn):
_seed(live_dsn, "src-a", TEXTS, _Embeddings(1.0))
before = _vectors(live_dsn, "src-a")
assert [row[0] for row in before] == TEXTS
assert all(row[1][0] == pytest.approx(1.0) for row in before)
new_model = _Embeddings(9.0)
with patch(
"application.vectorstore.base.BaseVectorStore._get_embeddings",
return_value=new_model,
):
seen, written = reembed.reembed_pgvector("src-a", batch_size=2, dry_run=False)
assert (seen, written) == (3, 3)
after = _vectors(live_dsn, "src-a")
# Text is untouched; only the vectors moved.
assert [row[0] for row in after] == TEXTS
assert all(row[1][0] == pytest.approx(9.0) for row in after)
def test_dry_run_counts_without_writing(self, live_dsn):
_seed(live_dsn, "src-b", TEXTS, _Embeddings(1.0))
new_model = _Embeddings(9.0)
with patch(
"application.vectorstore.base.BaseVectorStore._get_embeddings",
return_value=new_model,
):
seen, written = reembed.reembed_pgvector("src-b", batch_size=2, dry_run=True)
assert (seen, written) == (3, 0)
assert new_model.calls == 0, "dry run must not embed"
assert all(row[1][0] == pytest.approx(1.0) for row in _vectors(live_dsn, "src-b"))
def test_batches_are_respected(self, live_dsn):
_seed(live_dsn, "src-c", TEXTS, _Embeddings(1.0))
new_model = _Embeddings(9.0)
with patch(
"application.vectorstore.base.BaseVectorStore._get_embeddings",
return_value=new_model,
):
reembed.reembed_pgvector("src-c", batch_size=2, dry_run=False)
assert new_model.calls == 2, "3 chunks at batch 2 is two embed calls"
def test_only_the_named_source_is_touched(self, live_dsn):
_seed(live_dsn, "src-d", TEXTS, _Embeddings(1.0))
_seed(live_dsn, "src-e", TEXTS, _Embeddings(1.0))
with patch(
"application.vectorstore.base.BaseVectorStore._get_embeddings",
return_value=_Embeddings(9.0),
):
reembed.reembed_pgvector("src-d", batch_size=64, dry_run=False)
assert all(row[1][0] == pytest.approx(9.0) for row in _vectors(live_dsn, "src-d"))
assert all(row[1][0] == pytest.approx(1.0) for row in _vectors(live_dsn, "src-e"))
def test_source_discovery_lists_every_source(self, live_dsn):
_seed(live_dsn, "src-f", TEXTS, _Embeddings(1.0))
_seed(live_dsn, "src-g", TEXTS, _Embeddings(1.0))
assert reembed.list_source_ids("pgvector") == ["src-f", "src-g"]
GRAPH_SOURCE = "11111111-2222-3333-4444-555555555555"
def _seed_graph_node(dsn, source_id, name, seed):
"""Insert one graph node carrying a name embedding at ``seed``."""
from application.graphrag.store import GraphStore
store = PGVectorStore(source_id=source_id, connection_string=dsn)
conn = store._get_connection()
GraphStore.create_schema(conn, dimension=DIM)
cursor = conn.cursor()
try:
cursor.execute(
"INSERT INTO graph_nodes (id, source_id, name, normalized_name, type, "
"description, degree, doc_freq, name_embedding) "
"VALUES (gen_random_uuid(), %s, %s, %s, 'ENTITY', '', 0, 0, %s::vector)",
(source_id, name, name.lower(), str([seed] + [0.0] * (DIM - 1))),
)
conn.commit()
finally:
cursor.close()
store.close()
def _node_vectors(dsn, source_id):
store = PGVectorStore(source_id=source_id, connection_string=dsn)
conn = store._get_connection()
cursor = conn.cursor()
try:
cursor.execute(
"SELECT name, name_embedding FROM graph_nodes "
"WHERE source_id = %s ORDER BY name",
(source_id,),
)
return [(name, _as_list(vector)) for name, vector in cursor.fetchall()]
finally:
cursor.close()
store.close()
class TestGraphNodeReembedLive:
"""``graph_nodes.name_embedding`` seeds every graph traversal.
Rewriting only the chunk table leaves it in the previous model's space,
and because mpnet and granite-311m share a width the column accepts the
mismatch silently -- exactly the failure the script exists to prevent.
"""
def test_node_names_are_re_embedded(self, live_dsn, monkeypatch):
from application.core import settings as settings_module
monkeypatch.setattr(
settings_module.settings, "GRAPHRAG_ENABLED", True, raising=False
)
_seed(live_dsn, GRAPH_SOURCE, TEXTS, _Embeddings(1.0))
_seed_graph_node(live_dsn, GRAPH_SOURCE, "Alpha", 1.0)
assert all(v[0] == pytest.approx(1.0) for _, v in _node_vectors(live_dsn, GRAPH_SOURCE))
with patch(
"application.vectorstore.base.BaseVectorStore._get_embeddings",
return_value=_Embeddings(9.0),
):
reembed.reembed_pgvector(GRAPH_SOURCE, batch_size=64, dry_run=False)
after = _node_vectors(live_dsn, GRAPH_SOURCE)
assert [name for name, _ in after] == ["Alpha"]
assert all(v[0] == pytest.approx(9.0) for _, v in after), (
"graph node names must move with the chunk vectors"
)
def test_graph_is_left_alone_when_graphrag_is_off(self, live_dsn, monkeypatch):
from application.core import settings as settings_module
monkeypatch.setattr(
settings_module.settings, "GRAPHRAG_ENABLED", False, raising=False
)
_seed(live_dsn, GRAPH_SOURCE, TEXTS, _Embeddings(1.0))
_seed_graph_node(live_dsn, GRAPH_SOURCE, "Alpha", 1.0)
with patch(
"application.vectorstore.base.BaseVectorStore._get_embeddings",
return_value=_Embeddings(9.0),
):
reembed.reembed_pgvector(GRAPH_SOURCE, batch_size=64, dry_run=False)
assert all(v[0] == pytest.approx(1.0) for _, v in _node_vectors(live_dsn, GRAPH_SOURCE))
def test_dry_run_leaves_node_vectors_untouched(self, live_dsn, monkeypatch):
from application.core import settings as settings_module
monkeypatch.setattr(
settings_module.settings, "GRAPHRAG_ENABLED", True, raising=False
)
_seed(live_dsn, GRAPH_SOURCE, TEXTS, _Embeddings(1.0))
_seed_graph_node(live_dsn, GRAPH_SOURCE, "Alpha", 1.0)
with patch(
"application.vectorstore.base.BaseVectorStore._get_embeddings",
return_value=_Embeddings(9.0),
):
reembed.reembed_pgvector(GRAPH_SOURCE, batch_size=64, dry_run=True)
assert all(v[0] == pytest.approx(1.0) for _, v in _node_vectors(live_dsn, GRAPH_SOURCE))
@@ -290,3 +290,32 @@ class TestBucketedTotals:
)
assert len(rows) == 1
assert rows[0]["prompt_tokens"] == 10
class TestCacheBins:
def _row(self, pg_conn, user_id):
return pg_conn.execute(
text(
"SELECT cached_tokens, cache_write_tokens FROM token_usage "
"WHERE user_id = :u ORDER BY id DESC LIMIT 1"
),
{"u": user_id},
).fetchone()
def test_persists_cache_bins(self, pg_conn):
repo = TokenUsageRepository(pg_conn)
repo.insert(
user_id="cache-user",
prompt_tokens=1000,
generated_tokens=10,
cached_tokens=800,
cache_write_tokens=100,
)
assert tuple(self._row(pg_conn, "cache-user")) == (800, 100)
def test_cache_bins_default_to_null_not_zero(self, pg_conn):
"""NULL means "provider did not report"; 0 means "reported no cache
activity". Keeping them distinct is what makes a hit-rate query honest."""
repo = TokenUsageRepository(pg_conn)
repo.insert(user_id="cache-user-2", prompt_tokens=10, generated_tokens=1)
assert tuple(self._row(pg_conn, "cache-user-2")) == (None, None)
@@ -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,164 @@ 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
@pytest.mark.unit
class TestUnknownWidthIsProbed:
"""A remote server's width is only knowable by asking it.
``RemoteEmbeddings`` reports ``None`` until its first call, so sizing the
table from the attribute alone fell back to 768 and skipped the mismatch
check — the silent ``vector(768)`` column this hook exists to prevent.
"""
def _run(self, remote, table_dimension=1024):
conn = MagicMock()
conn.cursor.return_value = MagicMock()
with patch("psycopg.connect", return_value=conn), patch(
"application.vectorstore.model_registry.dimension_for", return_value=None
), patch(
"application.vectorstore.base.build_local_embeddings", return_value=remote
), patch(
"application.vectorstore.pgvector.PGVectorStore.create_schema"
) as vector_schema, patch(
"application.vectorstore.pgvector.PGVectorStore.table_dimension",
return_value=table_dimension,
):
try:
ensure_vector_schema()
raised = False
except RuntimeError:
raised = True
return vector_schema, raised
@staticmethod
def _remote(width=None, error=None):
remote = MagicMock()
remote.dimension = None
remote.embed_query.side_effect = error or (lambda _text: [0.0] * width)
return remote
def test_the_table_is_sized_from_the_probe(self, vector_settings):
remote = self._remote(width=1024)
vector_schema, _ = self._run(remote, table_dimension=1024)
remote.embed_query.assert_called_once()
assert vector_schema.call_args.kwargs["dimension"] == 1024
def test_the_probe_restores_the_mismatch_check(self, vector_settings):
_, raised = self._run(self._remote(width=768), table_dimension=1024)
assert raised, "a 768-dim model against a vector(1024) table must fail loudly"
def test_an_unreachable_server_does_not_block_boot(self, vector_settings):
vector_schema, raised = self._run(
self._remote(error=ConnectionError("server down")), table_dimension=1024
)
assert not raised
assert vector_schema.call_args.kwargs["dimension"] == 768
def test_a_model_that_knows_its_width_is_not_probed(self, vector_settings):
local = MagicMock()
local.dimension = 384
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"
+131
View File
@@ -0,0 +1,131 @@
"""Which embedding model an installation is pinned to."""
from unittest.mock import MagicMock, patch
import pytest
from application.storage.db import embeddings_pin
from application.storage.db.embeddings_pin import (
NOTICE_KEY,
PIN_KEY,
resolve_embeddings_pin,
)
from application.vectorstore.model_registry import DEFAULT_LEGACY, DEFAULT_NEW_INSTALL
@pytest.fixture
def store():
"""An in-memory stand-in for the ``app_metadata`` key/value table."""
data = {}
repo = MagicMock()
repo.get.side_effect = data.get
repo.set.side_effect = lambda k, v: data.__setitem__(k, v)
repo.setdefault.side_effect = lambda k, v: data.setdefault(k, v)
repo._data = data
return repo
def _run(store, *, has_sources, env_pinned=False, name="unset"):
session = MagicMock()
session.__enter__ = MagicMock(return_value=MagicMock())
session.__exit__ = MagicMock(return_value=False)
fields = {"EMBEDDINGS_NAME"} if env_pinned else set()
with patch.object(embeddings_pin.settings, "EMBEDDINGS_NAME", name), patch.object(
type(embeddings_pin.settings), "model_fields_set", property(lambda self: fields)
), patch.object(embeddings_pin, "db_session", return_value=session), patch.object(
embeddings_pin, "AppMetadataRepository", return_value=store
), patch.object(
embeddings_pin, "_has_sources", return_value=has_sources
):
resolve_embeddings_pin()
return embeddings_pin.settings.EMBEDDINGS_NAME
class TestFreshInstall:
def test_an_empty_installation_is_pinned_to_the_current_model(self, store):
assert _run(store, has_sources=False) == DEFAULT_NEW_INSTALL
assert store._data[PIN_KEY] == DEFAULT_NEW_INSTALL
def test_no_legacy_notice_is_printed(self, store, capsys):
_run(store, has_sources=False)
assert "reembed" not in capsys.readouterr().out
assert NOTICE_KEY not in store._data
class TestExistingInstall:
"""An index built by the old model must keep being read by it."""
def test_an_installation_with_sources_is_pinned_to_the_legacy_model(self, store):
assert _run(store, has_sources=True) == DEFAULT_LEGACY
assert store._data[PIN_KEY] == DEFAULT_LEGACY
def test_the_notice_names_the_migration_command(self, store, capsys):
_run(store, has_sources=True)
out = capsys.readouterr().out
assert "application.scripts.reembed" in out
assert DEFAULT_NEW_INSTALL in out
def test_the_notice_is_shown_only_once(self, store, capsys):
_run(store, has_sources=True)
capsys.readouterr()
store._data.pop(PIN_KEY) # force the decision again
_run(store, has_sources=True)
assert "reembed" not in capsys.readouterr().out
class TestPrecedence:
def test_a_stored_pin_survives_new_sources(self, store):
store._data[PIN_KEY] = DEFAULT_NEW_INSTALL
assert _run(store, has_sources=True) == DEFAULT_NEW_INSTALL
def test_the_environment_wins_and_nothing_is_stored(self, store):
assert _run(store, has_sources=True, env_pinned=True, name="my/model") == "my/model"
assert store._data == {}
def test_an_unreachable_database_leaves_the_default_alone(self):
with patch.object(embeddings_pin.settings, "EMBEDDINGS_NAME", "fallback"), patch.object(
type(embeddings_pin.settings), "model_fields_set", property(lambda self: set())
), patch.object(embeddings_pin, "db_session", side_effect=OSError("no db")):
resolve_embeddings_pin()
assert embeddings_pin.settings.EMBEDDINGS_NAME == "fallback"
class TestSourceModelMismatch:
"""The only signal that an index is being queried by the wrong model."""
def _run(self, rows, active, log):
conn = MagicMock()
conn.execute.side_effect = [
MagicMock(scalar=MagicMock(return_value="sources")),
MagicMock(fetchall=MagicMock(return_value=rows)),
]
session = MagicMock()
session.__enter__ = MagicMock(return_value=conn)
session.__exit__ = MagicMock(return_value=False)
with patch.object(embeddings_pin.settings, "EMBEDDINGS_NAME", active), patch.object(
embeddings_pin, "db_session", return_value=session
):
embeddings_pin.warn_on_source_model_mismatch(log)
def test_a_different_model_is_reported_with_counts(self):
log = MagicMock()
self._run([(DEFAULT_LEGACY, 28)], DEFAULT_NEW_INSTALL, log)
message = log.warning.call_args.args[0] % log.warning.call_args.args[1:]
assert "28 built with" in message
assert "application.scripts.reembed" in message
def test_an_alias_is_not_a_mismatch(self):
"""A stored alias and the canonical name are the same model."""
log = MagicMock()
self._run([("sentence-transformers/all-mpnet-base-v2", 28)], DEFAULT_LEGACY, log)
log.warning.assert_not_called()
def test_a_matching_model_says_nothing(self):
log = MagicMock()
self._run([(DEFAULT_NEW_INSTALL, 5)], DEFAULT_NEW_INSTALL, log)
log.warning.assert_not_called()
def test_two_unregistered_names_that_differ_are_a_mismatch(self):
log = MagicMock()
self._run([("some/other-model", 3)], "my/custom-model", log)
log.warning.assert_called_once()
+100
View File
@@ -0,0 +1,100 @@
"""Migration round-trip test for 0031_token_usage_cache_tokens."""
from __future__ import annotations
import os
import subprocess
import sys
from pathlib import Path
import pytest
from sqlalchemy import text
pytestmark = pytest.mark.integration
def _alembic_ini() -> Path:
return Path(__file__).resolve().parents[3] / "application" / "alembic.ini"
def _run_alembic(url: str, *args: str) -> None:
subprocess.check_call(
[sys.executable, "-m", "alembic", "-c", str(_alembic_ini()), *args],
timeout=60,
env={**os.environ, "POSTGRES_URI": url},
)
def _alembic_heads(url: str) -> list[str]:
out = subprocess.check_output(
[sys.executable, "-m", "alembic", "-c", str(_alembic_ini()), "heads"],
timeout=60,
env={**os.environ, "POSTGRES_URI": url},
text=True,
)
return [line for line in out.splitlines() if line.strip()]
def _alembic_version(conn) -> str:
return conn.execute(text("SELECT version_num FROM alembic_version")).scalar()
def _column_exists(conn, table: str, column: str) -> bool:
row = conn.execute(
text(
"SELECT 1 FROM information_schema.columns "
"WHERE table_name = :t AND column_name = :c AND table_schema = 'public'"
),
{"t": table, "c": column},
).fetchone()
return row is not None
_0031 = "0031_token_usage_cache_tokens"
_0030 = "0030_superseded_messages"
class TestMigration0031RoundTrip:
def test_single_head(self, pg_engine):
url = pg_engine.url.render_as_string(hide_password=False)
assert len(_alembic_heads(url)) == 1
def test_head_has_cache_columns(self, pg_engine):
with pg_engine.connect() as conn:
assert _alembic_version(conn) >= _0031
assert _column_exists(conn, "token_usage", "cached_tokens")
assert _column_exists(conn, "token_usage", "cache_write_tokens")
def test_downgrade_drops_then_upgrade_restores(self, pg_engine):
url = pg_engine.url.render_as_string(hide_password=False)
_run_alembic(url, "downgrade", _0030)
with pg_engine.connect() as conn:
assert _alembic_version(conn) == _0030
assert not _column_exists(conn, "token_usage", "cached_tokens")
assert not _column_exists(conn, "token_usage", "cache_write_tokens")
_run_alembic(url, "upgrade", "head")
with pg_engine.connect() as conn:
assert _alembic_version(conn) >= _0031
assert _column_exists(conn, "token_usage", "cached_tokens")
def test_existing_rows_read_null_cache_bins(self, pg_engine):
"""Rows written before 0031 must read NULL (unknown), not 0."""
url = pg_engine.url.render_as_string(hide_password=False)
_run_alembic(url, "downgrade", _0030)
with pg_engine.begin() as conn:
conn.execute(
text(
"INSERT INTO token_usage (user_id, prompt_tokens, generated_tokens) "
"VALUES ('u-mig31', 10, 1)"
)
)
_run_alembic(url, "upgrade", "head")
with pg_engine.connect() as conn:
row = conn.execute(
text(
"SELECT cached_tokens, cache_write_tokens FROM token_usage "
"WHERE user_id = 'u-mig31'"
)
).fetchone()
assert tuple(row) == (None, None)
+81 -37
View File
@@ -16,6 +16,13 @@ def local_storage(temp_base_dir):
return LocalStorage(base_dir=temp_base_dir)
@pytest.fixture
def real_storage(tmp_path):
"""Storage over a real directory, for the write paths worth exercising."""
base = os.path.realpath(str(tmp_path))
return LocalStorage(base_dir=base), base
@pytest.mark.unit
class TestLocalStorageInitialization:
@@ -37,50 +44,87 @@ class TestLocalStorageInitialization:
with pytest.raises(ValueError, match="Path traversal detected"):
local_storage._get_full_path("/absolute/path/test.txt")
@patch("os.makedirs")
@patch("builtins.open", new_callable=mock_open)
@patch("shutil.copyfileobj")
def test_save_file_creates_directory_and_saves(
self, mock_copyfileobj, mock_file, mock_makedirs, local_storage
):
file_data = io.BytesIO(b"test content")
path = "documents/test.txt"
def test_save_file_creates_directory_and_saves(self, real_storage):
storage, base = real_storage
result = storage.save_file(io.BytesIO(b"test content"), "documents/test.txt")
result = local_storage.save_file(file_data, path)
expected_dir = os.path.join(os.path.realpath("/tmp/test_storage"), "documents")
expected_file = os.path.join(os.path.realpath("/tmp/test_storage"), "documents/test.txt")
assert mock_makedirs.call_count == 1
assert os.path.normpath(mock_makedirs.call_args[0][0]) == os.path.normpath(
expected_dir
)
assert mock_makedirs.call_args[1] == {"exist_ok": True}
assert mock_file.call_count == 1
assert os.path.normpath(mock_file.call_args[0][0]) == os.path.normpath(
expected_file
)
assert mock_file.call_args[0][1] == "wb"
mock_copyfileobj.assert_called_once_with(file_data, mock_file())
written = os.path.join(base, "documents/test.txt")
assert os.path.isfile(written)
with open(written, "rb") as f:
assert f.read() == b"test content"
assert result == {"storage_type": "local"}
@patch("os.makedirs")
def test_save_file_with_save_method(self, mock_makedirs, local_storage):
file_data = MagicMock()
file_data.save = MagicMock()
path = "documents/test.txt"
def test_save_file_with_save_method(self, real_storage):
"""Werkzeug's ``FileStorage.save`` accepts the destination handle."""
storage, base = real_storage
result = local_storage.save_file(file_data, path)
class _Uploaded:
def save(self, dst):
dst.write(b"from save()")
expected_file = os.path.join(os.path.realpath("/tmp/test_storage"), "documents/test.txt")
assert file_data.save.call_count == 1
assert os.path.normpath(file_data.save.call_args[0][0]) == os.path.normpath(
expected_file
)
result = storage.save_file(_Uploaded(), "documents/test.txt")
with open(os.path.join(base, "documents/test.txt"), "rb") as f:
assert f.read() == b"from save()"
assert result == {"storage_type": "local"}
def test_save_file_replaces_existing_content(self, real_storage):
storage, base = real_storage
storage.save_file(io.BytesIO(b"first"), "a/b.bin")
storage.save_file(io.BytesIO(b"second"), "a/b.bin")
with open(os.path.join(base, "a/b.bin"), "rb") as f:
assert f.read() == b"second"
def test_failed_write_leaves_the_previous_file_intact(self, real_storage):
"""The whole point of writing through a temp file.
``reembed`` rewrites an index in place; a half-written ``index.faiss``
loads at neither the old width nor the new one, so an interrupted write
must leave the previous bytes untouched rather than truncate them.
"""
storage, base = real_storage
storage.save_file(io.BytesIO(b"the good index"), "indexes/s1/index.faiss")
class _DiesHalfway:
def __init__(self):
self._served = False
def read(self, size=-1):
if self._served:
raise OSError("connection reset")
self._served = True
return b"corrupt"
with pytest.raises(OSError, match="connection reset"):
storage.save_file(_DiesHalfway(), "indexes/s1/index.faiss")
with open(os.path.join(base, "indexes/s1/index.faiss"), "rb") as f:
assert f.read() == b"the good index"
def test_failed_write_leaves_no_temp_file_behind(self, real_storage):
storage, base = real_storage
class _Explodes:
def read(self, size=-1):
raise OSError("boom")
with pytest.raises(OSError):
storage.save_file(_Explodes(), "indexes/s1/index.faiss")
assert os.listdir(os.path.join(base, "indexes/s1")) == []
def test_save_file_keeps_the_existing_permissions(self, real_storage):
"""``mkstemp`` is 0600; the replacement must not silently narrow access."""
storage, base = real_storage
storage.save_file(io.BytesIO(b"one"), "a/b.bin")
target = os.path.join(base, "a/b.bin")
os.chmod(target, 0o640)
storage.save_file(io.BytesIO(b"two"), "a/b.bin")
assert os.stat(target).st_mode & 0o777 == 0o640
def test_save_file_with_absolute_path_outside_base_raises(self, local_storage):
file_data = io.BytesIO(b"test content")
path = "/absolute/path/test.txt"
+51 -1
View File
@@ -1,4 +1,4 @@
from unittest.mock import patch
from unittest.mock import MagicMock, patch
import pytest
from application.celery_init import make_celery
@@ -222,3 +222,53 @@ def test_unparseable_file_raises_the_non_retryable_type():
with pytest.raises(DocumentParseError, match="No text could be extracted"):
embed_and_store_documents([], "/tmp", "src", None)
@pytest.mark.unit
class TestReclaimIsSkippedForEmbeds:
"""The post-task heap reclaim must not run on the query hot path.
``_reclaim_memory_after_task`` exists for the large transient allocations
docling/torch parsing makes. Query embedding became a Celery task, and a
full generational collect on a worker holding the ONNX model measured ~86 ms
against ~8 ms for the embed itself -- a 9x slowdown of the round trip for a
task that allocates a few kilobytes.
"""
@staticmethod
def _collects(task_name):
from application.celery_init import _reclaim_memory_after_task
task = MagicMock()
task.name = task_name
with patch("application.celery_init.gc.collect") as collect, patch(
"application.celery_init._trim_native_heap"
):
_reclaim_memory_after_task(task=task, task_id="t", state="SUCCESS")
return collect.called
def test_the_embed_task_is_skipped(self):
assert not self._collects("application.vectorstore.embeddings_tasks.embed_texts")
def test_parsing_still_reclaims(self):
assert self._collects("application.api.user.tasks.parse_document")
def test_ingest_still_reclaims(self):
assert self._collects("application.api.user.tasks.ingest")
def test_an_unnamed_sender_still_reclaims(self):
"""Unknown callers keep the old behaviour rather than silently skipping."""
from application.celery_init import _reclaim_memory_after_task
with patch("application.celery_init.gc.collect") as collect, patch(
"application.celery_init._trim_native_heap"
):
_reclaim_memory_after_task(task_id="t", state="SUCCESS")
assert collect.called
def test_the_skip_list_names_the_real_task(self):
"""A renamed task must not silently start paying the collect again."""
from application.celery_init import _NO_RECLAIM_TASKS
from application.vectorstore.embeddings_delegated import EMBED_TASK
assert EMBED_TASK in _NO_RECLAIM_TASKS
+220
View File
@@ -1,5 +1,6 @@
"""Tests for bounded user-upload stream helpers."""
import codecs
import io
import pytest
@@ -37,3 +38,222 @@ def test_limited_text_read_decodes_incrementally_and_rejects_overflow():
with pytest.raises(UploadTooLargeError):
read_text_upload_limited(_ShortReadStream(b"12345"), max_bytes=4)
# --- attachment type gate -------------------------------------------------
#
# ``SimpleDirectoryReader`` falls through to a plain-text ``open()`` for any
# suffix without a parser. That is what a .py or a .log attachment relies on,
# and it is also how a phone-uploaded video used to be "parsed" into
# megabytes of binary garbage, truncated, and stored with
# ``extraction.status == "ok"``. So a suffix with no parser is admitted on
# content: text in, binary out.
MP4_HEADER = b"\x00\x00\x00\x18ftypisom\x00\x00\x02\x00isomiso2avc1mp41"
@pytest.mark.parametrize(
("filename", "content"),
[
("clip.mp4", MP4_HEADER),
("movie.MOV", MP4_HEADER),
("archive.zip", b"PK\x03\x04\x14\x00\x00\x00\x08\x00" + bytes(range(32))),
("binary", bytes(range(32)) * 8),
("trailing.", b"\x00\x01\x02\x03"),
("x.tar.gz", b"\x1f\x8b\x08\x00\x00\x00\x00\x00\x00\x03"),
],
)
def test_enforce_parseable_attachment_rejects_binary_without_a_parser(
filename, content, tmp_path
):
from application.upload_limits import (
enforce_parseable_attachment,
UnsupportedUploadTypeError,
unsupported_upload_message,
)
path = tmp_path / "staged.bin"
path.write_bytes(content)
with pytest.raises(UnsupportedUploadTypeError) as excinfo:
enforce_parseable_attachment(path, filename)
assert str(excinfo.value) == unsupported_upload_message(filename)
assert str(excinfo.value).startswith("Unsupported file type")
def test_enforce_parseable_attachment_rejects_binary_named_as_text(tmp_path):
""".txt has no parser — it *is* the fallthrough — so it is sniffed like any suffix.
Renaming a video to notes.txt would otherwise walk straight back into the
bug this gate exists for.
"""
from application.upload_limits import (
enforce_parseable_attachment,
UnsupportedUploadTypeError,
)
path = tmp_path / "notes.txt"
path.write_bytes(MP4_HEADER + bytes(range(256)) * 4)
with pytest.raises(UnsupportedUploadTypeError) as excinfo:
enforce_parseable_attachment(path, "notes.txt")
assert str(excinfo.value) == "Unsupported file type: .txt"
@pytest.mark.parametrize(
"content",
[
codecs.BOM_UTF8 + "hello — Unicode\n".encode(),
codecs.BOM_UTF16_LE + "hello\n".encode("utf-16-le"),
codecs.BOM_UTF16_BE + "hello\n".encode("utf-16-be"),
codecs.BOM_UTF32_BE + "hello\n".encode("utf-32-be"),
],
)
def test_enforce_parseable_attachment_accepts_bom_marked_unicode_text(
content, tmp_path
):
"""A UTF-16 .txt is half NUL bytes and still ordinary text — the BOM says so."""
from application.upload_limits import enforce_parseable_attachment
path = tmp_path / "notes.txt"
path.write_bytes(content)
enforce_parseable_attachment(path, "notes.txt")
@pytest.mark.parametrize(
"bom",
[codecs.BOM_UTF8, codecs.BOM_UTF16_LE, codecs.BOM_UTF16_BE, codecs.BOM_UTF32_BE],
)
def test_enforce_parseable_attachment_rejects_binary_behind_a_bom(bom, tmp_path):
"""A BOM says which encoding to read, not that the content is text.
Otherwise three prepended bytes buy any binary a pass.
"""
from application.upload_limits import (
enforce_parseable_attachment,
UnsupportedUploadTypeError,
)
path = tmp_path / "notes.txt"
path.write_bytes(bom + MP4_HEADER + bytes(range(256)) * 8)
with pytest.raises(UnsupportedUploadTypeError):
enforce_parseable_attachment(path, "notes.txt")
def test_enforce_parseable_attachment_uses_the_extractor_it_is_given(tmp_path):
"""The worker holds the live parser table; a trimmed install must not admit on trust.
Without docling the fallback extractor has no .webp handler, so a .webp
would otherwise skip the content check and be read as plain text.
"""
from application.upload_limits import (
enforce_parseable_attachment,
UnsupportedUploadTypeError,
)
path = tmp_path / "scan.webp"
path.write_bytes(b"RIFF\x00\x00\x00\x00WEBPVP8 " + bytes(range(256)))
# Default list: .webp is parser-backed, admitted on its name.
enforce_parseable_attachment(path, "scan.webp")
# The extractor actually loaded has no .webp parser.
with pytest.raises(UnsupportedUploadTypeError):
enforce_parseable_attachment(path, "scan.webp", {".pdf", ".docx"})
@pytest.mark.parametrize(
"filename",
[
"Report.PDF",
"photo.JPG",
"slides.pptx",
"voice.ogg",
"page.xhtml",
"doc.adoc",
"scan.webp",
"fax.tiff",
"subs.vtt",
"feed.xml",
],
)
def test_enforce_parseable_attachment_accepts_parser_backed_types(filename, tmp_path):
"""A parser-backed suffix is admitted on its name — a PDF is binary and parses fine."""
from application.upload_limits import enforce_parseable_attachment
path = tmp_path / "staged.bin"
path.write_bytes(MP4_HEADER)
enforce_parseable_attachment(path, filename)
@pytest.mark.parametrize(
"filename",
[
"notes.txt",
"main.py",
"server.log",
"config.yaml",
"query.sql",
"Dockerfile",
"notes.unknown",
],
)
def test_enforce_parseable_attachment_accepts_text_without_a_parser(filename, tmp_path):
"""The plain-text fallthrough reads these correctly, so they must stay allowed."""
from application.upload_limits import enforce_parseable_attachment
path = tmp_path / "staged.txt"
path.write_text("def main():\n\treturn 'café — ok'\n", encoding="utf-8")
enforce_parseable_attachment(path, filename)
@pytest.mark.parametrize(
("sample", "expected"),
[
(b"", True),
(b"plain text\n", True),
("café — em dash\n".encode(), True),
(b"\x1b[31mred log line\x1b[0m\n", True),
(codecs.BOM_UTF16_LE + "hi\n".encode("utf-16-le"), True),
(codecs.BOM_UTF8 + b"hi\n", True),
(codecs.BOM_UTF32_LE + "hi\n".encode("utf-32-le"), True),
# A BOM in front of binary is still binary.
(codecs.BOM_UTF8 + b"\x00\x01\x02", False),
(codecs.BOM_UTF16_LE + MP4_HEADER, False),
(b"text\x00with nul", False),
(bytes(range(32)) * 4, False),
(b"\x7f\x7f\x7f\x7f" + b"a" * 16, False),
],
)
def test_looks_like_text(sample, expected):
from application.upload_limits import looks_like_text
assert looks_like_text(sample) is expected
def test_file_looks_like_text_only_samples_the_head(tmp_path):
"""Binary past the sampled head is the parser's problem, not the gate's."""
from application.upload_limits import file_looks_like_text
path = tmp_path / "staged.log"
path.write_bytes(b"a" * 9000 + b"\x00" * 100)
assert file_looks_like_text(path) is True
def test_file_looks_like_text_allows_an_unreadable_file(tmp_path):
from application.upload_limits import file_looks_like_text
assert file_looks_like_text(tmp_path / "missing.txt") is True
def test_unsupported_upload_message_names_the_extension():
from application.upload_limits import unsupported_upload_message
assert unsupported_upload_message("clip.mp4") == "Unsupported file type: .mp4"
assert unsupported_upload_message("Clip.MP4") == "Unsupported file type: .mp4"
assert unsupported_upload_message("binary") == "Unsupported file type: (no extension)"
+112
View File
@@ -600,3 +600,115 @@ class TestCountPromptTokens:
]
tokens = _count_prompt_tokens(messages, tools=None)
assert tokens > 0
# ── prompt-cache breakdown ────────────────────────────────────────────────────
#
# Providers report ``cached_tokens`` (and, on newer OpenAI-family models,
# ``cache_write_tokens``) as breakdowns of ``prompt_tokens``. The prompt bin
# stays the provider total (never subtract the details back out); the two
# sub-bins ride alongside so persistence and the finish events can chart a
# cache hit rate.
class _ReportingLLM:
decoded_token = {"sub": "user_1"}
user_api_key = None
agent_id = None
token_usage = {"prompt_tokens": 0, "generated_tokens": 0}
def __init__(self, details):
self._last_usage = {
"prompt_tokens": 1000,
"completion_tokens": 10,
"total_tokens": 1010,
}
if details is not None:
self._last_usage["prompt_tokens_details"] = details
self._last_usage_claimed = False
self.emitted = []
def _emit_gen_finished_log(self, model, **kwargs):
self.emitted.append(kwargs)
@pytest.mark.unit
def test_prefer_provider_usage_carries_cache_bins():
from application.usage import _prefer_provider_usage
llm = _ReportingLLM({"cached_tokens": 800, "cache_write_tokens": 100})
usage = _prefer_provider_usage(llm, {"prompt_tokens": 1, "generated_tokens": 1})
assert usage["prompt_tokens"] == 1000
assert usage["generated_tokens"] == 10
assert usage["cached_tokens"] == 800
assert usage["cache_write_tokens"] == 100
@pytest.mark.unit
def test_prefer_provider_usage_maps_anthropic_cache_creation_to_writes():
from application.usage import _prefer_provider_usage
llm = _ReportingLLM({"cached_tokens": 3, "cache_creation_tokens": 9})
usage = _prefer_provider_usage(llm, {"prompt_tokens": 1, "generated_tokens": 1})
assert usage["cached_tokens"] == 3
assert usage["cache_write_tokens"] == 9
@pytest.mark.unit
def test_prefer_provider_usage_without_details_has_no_cache_keys():
from application.usage import _prefer_provider_usage
llm = _ReportingLLM(None)
usage = _prefer_provider_usage(llm, {"prompt_tokens": 1, "generated_tokens": 1})
assert "cached_tokens" not in usage
assert "cache_write_tokens" not in usage
@pytest.mark.unit
def test_prefer_provider_usage_keeps_the_rest_of_the_call_record():
from application.usage import _prefer_provider_usage
llm = _ReportingLLM(None)
usage = _prefer_provider_usage(
llm, {"prompt_tokens": 1, "generated_tokens": 1, "model": "m"}
)
assert usage["model"] == "m"
@pytest.mark.unit
def test_decorator_persists_and_emits_cache_bins(monkeypatch):
_install_fake_token_repo(monkeypatch)
llm = _ReportingLLM({"cached_tokens": 800, "cache_write_tokens": 100})
@gen_token_usage
def wrapped(self, model, messages, stream, tools, **kwargs):
_ = (model, messages, stream, tools, kwargs)
return "ok"
wrapped(llm, "m", [{"role": "user", "content": "hi"}], False, None)
row = _FakeTokenUsageRepo.last_instance.inserted[0]
assert row["prompt_tokens"] == 1000
assert row["cached_tokens"] == 800
assert row["cache_write_tokens"] == 100
assert llm.emitted[0]["cached_tokens"] == 800
assert llm.emitted[0]["cache_write_tokens"] == 100
@pytest.mark.unit
def test_decorator_omits_cache_bins_when_provider_reports_none(monkeypatch):
_install_fake_token_repo(monkeypatch)
llm = _ReportingLLM(None)
@gen_token_usage
def wrapped(self, model, messages, stream, tools, **kwargs):
_ = (model, messages, stream, tools, kwargs)
return "ok"
wrapped(llm, "m", [{"role": "user", "content": "hi"}], False, None)
row = _FakeTokenUsageRepo.last_instance.inserted[0]
assert row["cached_tokens"] is None
assert row["cache_write_tokens"] is None
assert llm.emitted[0]["cached_tokens"] is None
assert llm.emitted[0]["cache_write_tokens"] is None
+68 -23
View File
@@ -263,11 +263,11 @@ class TestEmbeddingsSingleton:
def test_get_instance_hf_ignores_positional_key(
self, mock_get_wrapper, mock_settings
):
"""A stray key must not reach the zero-arg HuggingFace factory.
"""A stray key must not reach the wrapper for a registered model.
The factories are ``lambda: EmbeddingsWrapper(...)``, so a caller that
passed ``settings.EMBEDDINGS_KEY`` positionally used to blow up with
``TypeError: <lambda>() takes 0 positional arguments``.
Registered models take their whole configuration from the registry, so
a caller that passes ``settings.EMBEDDINGS_KEY`` positionally (as the
vector stores do) must have it dropped rather than forwarded.
"""
mock_settings.EMBEDDINGS_BASE_URL = None
mock_wrapper_cls = Mock()
@@ -278,9 +278,8 @@ class TestEmbeddingsSingleton:
result = EmbeddingsSingleton.get_instance(HF_MPNET, None)
assert result is mock_instance
mock_wrapper_cls.assert_called_once_with(
"sentence-transformers/all-mpnet-base-v2"
)
# The configured name is passed through; the registry maps it to a repo.
mock_wrapper_cls.assert_called_once_with(HF_MPNET)
@patch("application.vectorstore.base.settings")
@patch("application.vectorstore.base._get_embeddings_wrapper")
@@ -293,9 +292,7 @@ class TestEmbeddingsSingleton:
EmbeddingsSingleton.get_instance(HF_MPNET, openai_api_key="sk-nope")
mock_wrapper_cls.assert_called_once_with(
"sentence-transformers/all-mpnet-base-v2"
)
mock_wrapper_cls.assert_called_once_with(HF_MPNET)
# --- BaseVectorStore ---
@@ -401,12 +398,15 @@ class TestBaseVectorStore:
@patch("application.vectorstore.base.settings")
@patch("application.vectorstore.base.EmbeddingsSingleton.get_instance")
@patch("os.path.exists")
def test_get_embeddings_huggingface_local_model(
self, mock_exists, mock_get_instance, mock_settings
def test_get_embeddings_registered_model_passes_configured_name(
self, mock_get_instance, mock_settings
):
"""No bundled-path branch any more: the name goes straight through.
FastEmbed resolves artifacts through its own cache (warmed in the
image), so the old ``/app/models/...`` probe has no job to do.
"""
mock_settings.EMBEDDINGS_BASE_URL = None
mock_exists.side_effect = lambda p: p == "/app/models/all-mpnet-base-v2"
mock_emb = Mock()
mock_get_instance.return_value = mock_emb
@@ -415,7 +415,9 @@ class TestBaseVectorStore:
"huggingface_sentence-transformers/all-mpnet-base-v2"
)
assert result is mock_emb
mock_get_instance.assert_called_with("/app/models/all-mpnet-base-v2")
mock_get_instance.assert_called_with(
"huggingface_sentence-transformers/all-mpnet-base-v2"
)
@patch("application.vectorstore.base.settings")
@patch("application.vectorstore.base.EmbeddingsSingleton.get_instance")
@@ -518,18 +520,18 @@ class TestGetEmbeddingsResolver:
@patch("application.vectorstore.base.settings")
@patch("application.vectorstore.base._get_embeddings_wrapper")
@patch("os.path.exists")
def test_uses_local_model_path_and_caches_it(
self, mock_exists, mock_get_wrapper, mock_settings
def test_repeated_resolution_loads_one_model(
self, mock_get_wrapper, mock_settings
):
"""With the bundled model present, the instance is keyed by its path.
"""A second call must not load a second copy of the model.
A second call must not load a second copy of the model.
The instance is keyed by the configured name. It used to be keyed by a
bundled filesystem path when one happened to exist, which meant the
same model could be cached twice under two keys.
"""
mock_settings.EMBEDDINGS_BASE_URL = None
mock_settings.EMBEDDINGS_NAME = HF_MPNET
mock_settings.EMBEDDINGS_KEY = None
mock_exists.side_effect = lambda path: path == LOCAL_MPNET
mock_wrapper_cls = Mock()
mock_wrapper_cls.return_value = Mock()
mock_get_wrapper.return_value = mock_wrapper_cls
@@ -538,8 +540,8 @@ class TestGetEmbeddingsResolver:
second = get_embeddings()
assert first is second
assert set(EmbeddingsSingleton._instances) == {LOCAL_MPNET}
mock_wrapper_cls.assert_called_once_with(LOCAL_MPNET)
assert set(EmbeddingsSingleton._instances) == {HF_MPNET}
mock_wrapper_cls.assert_called_once_with(HF_MPNET)
@patch("application.vectorstore.base.settings")
def test_remote_when_base_url_configured(self, mock_settings):
@@ -589,6 +591,49 @@ class TestGetEmbeddingsResolver:
"openai_text-embedding-ada-002", model="embed-deploy"
)
@patch("application.vectorstore.base.settings")
@patch("application.vectorstore.base.EmbeddingsSingleton.get_instance")
def test_openai_alias_also_reaches_the_azure_deployment(
self, mock_get_instance, mock_settings
):
"""The registry accepts the bare alias, so the key handling must too.
Matching on the canonical string alone sent the alias down the generic
branch, where the deployment name is never passed and Azure answers
every embed with DeploymentNotFound.
"""
mock_settings.EMBEDDINGS_BASE_URL = None
mock_settings.OPENAI_API_BASE = "https://azure.openai.com"
mock_settings.OPENAI_API_VERSION = "2023-05-15"
mock_settings.AZURE_DEPLOYMENT_NAME = "deploy"
mock_settings.AZURE_EMBEDDINGS_DEPLOYMENT_NAME = "embed-deploy"
mock_settings.EMBEDDINGS_NAME = "text-embedding-ada-002"
mock_settings.EMBEDDINGS_KEY = "sk-key"
get_embeddings()
mock_get_instance.assert_called_once_with(
"text-embedding-ada-002", model="embed-deploy"
)
@patch("application.vectorstore.base.settings")
@patch("application.vectorstore.base.EmbeddingsSingleton.get_instance")
def test_openai_name_is_matched_case_insensitively(
self, mock_get_instance, mock_settings
):
mock_settings.EMBEDDINGS_BASE_URL = None
mock_settings.OPENAI_API_BASE = None
mock_settings.OPENAI_API_VERSION = None
mock_settings.AZURE_DEPLOYMENT_NAME = None
mock_settings.EMBEDDINGS_NAME = "OpenAI_Text-Embedding-Ada-002"
mock_settings.EMBEDDINGS_KEY = "sk-from-settings"
get_embeddings()
mock_get_instance.assert_called_once_with(
"OpenAI_Text-Embedding-Ada-002", openai_api_key="sk-from-settings"
)
@patch("application.vectorstore.base.settings")
@patch("application.vectorstore.base.EmbeddingsSingleton.get_instance")
def test_explicit_arguments_win_over_settings(
@@ -0,0 +1,380 @@
"""Query embedding runs on the worker so the API holds no model."""
import threading
import time
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")
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 TestTheConcurrentFirstWave:
"""The latch cannot cover requests already in flight beside the first one.
Nothing is latched until that first ``get()`` returns, so every thread in
the opening wave would otherwise block for the full
``EMBEDDINGS_DELEGATE_TIMEOUT`` at once -- at the shipped 60s across a 96
thread WSGI pool, an API that serves nothing at all.
"""
@staticmethod
def _blocking_celery(release, outcome):
"""A worker whose ``get`` blocks until ``release`` is set."""
def get(timeout=None):
release.wait(5)
if isinstance(outcome, Exception):
raise outcome
return outcome
result = MagicMock()
result.get.side_effect = get
celery = MagicMock()
celery.send_task.return_value = result
return celery
def _race(self, celery, embeddings, release, threads=8):
"""Start ``threads`` embeds, let them pile up, then unblock the prober."""
errors, values = [], []
started = threading.Barrier(threads + 1)
def call():
started.wait(5)
try:
values.append(embeddings.embed_query("q"))
except Exception as exc: # noqa: BLE001 -- recorded for the assertions
errors.append(exc)
with patch.dict("sys.modules", {"application.celery_init": MagicMock(celery=celery)}):
workers = [threading.Thread(target=call) for _ in range(threads)]
for worker in workers:
worker.start()
started.wait(5)
time.sleep(0.1) # let the followers reach the probe gate
release.set()
for worker in workers:
worker.join(10)
return values, errors
def test_only_one_caller_waits_on_an_unproven_worker(self, not_in_worker):
release = threading.Event()
celery = self._blocking_celery(release, TimeoutError("no worker"))
embeddings = DelegatedEmbeddings("granite-311m")
with patch(
"application.vectorstore.embeddings_delegated._PROBE_WAIT", 0.05
):
values, errors = self._race(celery, embeddings, release)
assert values == []
assert len(errors) == 8
# One probe published; the rest gave up without their own round trip.
assert celery.send_task.call_count == 1
assert sum("still unanswered" in str(e) for e in errors) == 7
def test_the_fast_failure_still_names_the_remedy(self, not_in_worker):
release = threading.Event()
celery = self._blocking_celery(release, TimeoutError("no worker"))
embeddings = DelegatedEmbeddings("granite-311m")
with patch("application.vectorstore.embeddings_delegated._PROBE_WAIT", 0.05):
_, errors = self._race(celery, embeddings, release, threads=3)
assert all("EMBEDDINGS_DELEGATE_TO_WORKER=false" in str(e) for e in errors)
def test_a_healthy_worker_serves_the_whole_wave(self, not_in_worker):
release = threading.Event()
celery = self._blocking_celery(release, [[0.1, 0.2]])
embeddings = DelegatedEmbeddings("granite-311m")
values, errors = self._race(celery, embeddings, release)
assert errors == []
assert values == [[0.1, 0.2]] * 8
# The probe proves the worker, then every follower dispatches for real.
assert celery.send_task.call_count == 8
assert embeddings._verified is True
def test_a_proven_worker_adds_no_gate(self, not_in_worker):
"""After one success the probe is out of the path entirely."""
release = threading.Event()
release.set()
celery = self._blocking_celery(release, [[0.3]])
embeddings = DelegatedEmbeddings("granite-311m")
with patch.dict("sys.modules", {"application.celery_init": MagicMock(celery=celery)}):
embeddings.embed_query("warm")
assert embeddings._verified is True
with patch.object(embeddings, "_state_lock") as lock:
with patch.dict(
"sys.modules", {"application.celery_init": MagicMock(celery=celery)}
):
embeddings.embed_query("q")
lock.__enter__.assert_not_called()
def test_a_proven_worker_that_dies_is_gated_again(self, not_in_worker):
"""The proof must not outlive the worker that supplied it.
A worker that is redeployed or OOM-killed is the failure that actually
happens in production, and it is the one the probe gate stopped
covering: ``_verified`` short-circuits ahead of it. The wave in flight
when the worker dies cannot be saved -- every caller is already past
the check -- but every wave after it must be gated again.
"""
warm = threading.Event()
warm.set()
healthy = self._blocking_celery(warm, [[0.4]])
embeddings = DelegatedEmbeddings("granite-311m")
with patch.dict(
"sys.modules", {"application.celery_init": MagicMock(celery=healthy)}
):
embeddings.embed_query("warm")
assert embeddings._verified is True
# No cooldown, so anything that gates the second wave can only be the
# probe -- which engages only because the failure cleared _verified.
with patch(
"application.vectorstore.embeddings_delegated._FAILURE_COOLDOWN", 0.0
), patch("application.vectorstore.embeddings_delegated._PROBE_WAIT", 0.05):
dying = threading.Event()
died = self._blocking_celery(dying, TimeoutError("worker went away"))
self._race(died, embeddings, dying)
# The wave that was already in flight all dispatched, as it must.
assert died.send_task.call_count == 8
assert embeddings._verified is False
again = threading.Event()
still_dead = self._blocking_celery(again, TimeoutError("still gone"))
values, errors = self._race(still_dead, embeddings, again)
assert values == []
assert len(errors) == 8
# One probe pays the timeout; the other seven fail fast.
assert still_dead.send_task.call_count == 1
assert sum("still unanswered" in str(e) for e in errors) == 7
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")
+350 -113
View File
@@ -1,156 +1,393 @@
from unittest.mock import MagicMock, Mock, patch
"""Local embeddings run through FastEmbed, configured from the model registry."""
from unittest.mock import MagicMock, patch
import numpy as np
import pytest
from application.vectorstore import embeddings_local
from application.vectorstore.embeddings_local import EmbeddingsWrapper
from application.vectorstore.model_registry import GRANITE_97M, MPNET
@pytest.mark.unit
class TestEmbeddingsWrapper:
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_init_success(self, mock_st_cls):
mock_model = MagicMock()
mock_model._first_module.return_value = MagicMock()
mock_model.get_embedding_dimension.return_value = 768
mock_st_cls.return_value = mock_model
from application.vectorstore.embeddings_local import EmbeddingsWrapper
@pytest.fixture(autouse=True)
def _clear_registration():
"""``add_custom_model`` writes to a FastEmbed global; keep tests isolated."""
embeddings_local._registered.clear()
yield
embeddings_local._registered.clear()
wrapper = EmbeddingsWrapper("test-model")
mock_st_cls.assert_called_once()
assert wrapper.dimension == 768
@pytest.fixture(autouse=True)
def _no_hub_reads():
"""Keep unit tests off the network.
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_init_falls_back_to_legacy_dimension_api(self, mock_st_cls):
"""sentence-transformers < 5.4 only exposes get_sentence_embedding_dimension."""
mock_model = MagicMock()
mock_model._first_module.return_value = MagicMock()
del mock_model.get_embedding_dimension
mock_model.get_sentence_embedding_dimension.return_value = 768
mock_st_cls.return_value = mock_model
``_spec_for`` now asks a repository how it pools; without this every test
naming an unregistered model would reach the Hugging Face hub. ``None`` is
the "declares nothing" answer, which is the behaviour these tests were
written against. Tests that exercise the metadata patch it themselves.
"""
with patch.object(embeddings_local, "_read_repo_json", return_value=None):
yield
from application.vectorstore.embeddings_local import EmbeddingsWrapper
wrapper = EmbeddingsWrapper("test-model")
@pytest.fixture
def fake_fastembed():
"""Patch FastEmbed so no model is downloaded or run."""
text_embedding = MagicMock()
instance = MagicMock()
instance.embed.return_value = iter([np.array([0.1, 0.2, 0.3])])
text_embedding.return_value = instance
# Registration checks this before calling ``add_custom_model``; an empty
# list means "no built-in collides", which is the case for every name in
# our registry.
text_embedding.list_supported_models.return_value = []
with patch("fastembed.TextEmbedding", text_embedding):
yield text_embedding, instance
assert wrapper.dimension == 768
mock_model.get_sentence_embedding_dimension.assert_called_once()
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_init_failure(self, mock_st_cls):
mock_st_cls.side_effect = Exception("model not found")
class TestBuiltinModelRegistration:
"""FastEmbed ships ~30 models of its own and refuses to re-register any of
them, so registering unconditionally broke every natively-supported name."""
from application.vectorstore.embeddings_local import EmbeddingsWrapper
def test_builtin_name_is_not_re_registered(self, fake_fastembed):
text_embedding, _ = fake_fastembed
text_embedding.list_supported_models.return_value = [
{"model": "BAAI/bge-small-en-v1.5"}
]
EmbeddingsWrapper("BAAI/bge-small-en-v1.5")
text_embedding.add_custom_model.assert_not_called()
assert text_embedding.call_args.kwargs["model_name"] == "BAAI/bge-small-en-v1.5"
with pytest.raises(Exception, match="model not found"):
EmbeddingsWrapper("bad-model")
def test_builtin_match_ignores_case(self, fake_fastembed):
text_embedding, _ = fake_fastembed
text_embedding.list_supported_models.return_value = [
{"model": "baai/BGE-Small-EN-v1.5"}
]
EmbeddingsWrapper("BAAI/bge-small-en-v1.5")
text_embedding.add_custom_model.assert_not_called()
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_init_none_model(self, mock_st_cls):
mock_st_cls.return_value = None
def test_unknown_name_is_still_registered(self, fake_fastembed):
text_embedding, _ = fake_fastembed
text_embedding.list_supported_models.return_value = [
{"model": "BAAI/bge-small-en-v1.5"}
]
EmbeddingsWrapper("some-org/custom-embedder")
text_embedding.add_custom_model.assert_called_once()
from application.vectorstore.embeddings_local import EmbeddingsWrapper
def test_real_fastembed_accepts_its_own_builtin(self):
"""Runs against the installed FastEmbed, not the MagicMock.
with pytest.raises((ValueError, AttributeError)):
EmbeddingsWrapper("bad-model")
The mocked tests above cannot catch this: the failure was
``add_custom_model`` raising, and a MagicMock never raises.
"""
fastembed = pytest.importorskip("fastembed")
builtins = [m["model"] for m in fastembed.TextEmbedding.list_supported_models()]
assert builtins, "expected FastEmbed to ship built-in models"
spec = embeddings_local._spec_for(builtins[0])
# Must not raise ValueError("... is already registered ...").
embeddings_local._register(spec)
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_init_null_first_module(self, mock_st_cls):
mock_model = MagicMock()
mock_model._first_module.return_value = None
mock_st_cls.return_value = mock_model
from application.vectorstore.embeddings_local import EmbeddingsWrapper
class TestRegistryDrivenLoading:
def test_registered_model_loads_by_repo_not_by_configured_name(self, fake_fastembed):
text_embedding, _ = fake_fastembed
wrapper = EmbeddingsWrapper(MPNET.name)
assert text_embedding.call_args.kwargs["model_name"] == MPNET.repo
assert wrapper.dimension == MPNET.dimension
with pytest.raises(ValueError, match="failed to load properly"):
EmbeddingsWrapper("bad-model")
def test_legacy_alias_resolves_to_the_same_model(self, fake_fastembed):
text_embedding, _ = fake_fastembed
EmbeddingsWrapper("huggingface_sentence-transformers-all-mpnet-base-v2")
assert text_embedding.call_args.kwargs["model_name"] == MPNET.repo
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_embed_query(self, mock_st_cls):
mock_model = MagicMock()
mock_model._first_module.return_value = MagicMock()
mock_model.get_sentence_embedding_dimension.return_value = 3
mock_model.encode.return_value = MagicMock(tolist=Mock(return_value=[0.1, 0.2, 0.3]))
mock_st_cls.return_value = mock_model
def test_dimension_comes_from_registry_without_running_the_model(self, fake_fastembed):
_, instance = fake_fastembed
wrapper = EmbeddingsWrapper(GRANITE_97M.name)
assert wrapper.dimension == 384
instance.embed.assert_not_called()
from application.vectorstore.embeddings_local import EmbeddingsWrapper
def test_unknown_model_is_treated_as_a_hf_repo(self, fake_fastembed):
text_embedding, _ = fake_fastembed
wrapper = EmbeddingsWrapper("some-org/custom-embedder")
assert text_embedding.call_args.kwargs["model_name"] == "some-org/custom-embedder"
# No registry entry means no known width, so it must be probed.
assert wrapper.dimension == 3
wrapper = EmbeddingsWrapper("model")
result = wrapper.embed_query("hello world")
def test_load_failure_names_the_model_and_the_known_ones(self):
with patch("fastembed.TextEmbedding", side_effect=OSError("no such repo")):
with pytest.raises(RuntimeError) as excinfo:
EmbeddingsWrapper("broken/model")
message = str(excinfo.value)
assert "broken/model" in message
assert MPNET.name in message
mock_model.encode.assert_called_once_with("hello world")
assert result == [0.1, 0.2, 0.3]
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_embed_documents(self, mock_st_cls):
mock_model = MagicMock()
mock_model._first_module.return_value = MagicMock()
mock_model.get_sentence_embedding_dimension.return_value = 3
mock_model.encode.return_value = MagicMock(
tolist=Mock(return_value=[[0.1, 0.2], [0.3, 0.4]])
class TestSettingsPassthrough:
def test_threads_forwarded_when_configured(self, fake_fastembed):
text_embedding, _ = fake_fastembed
with patch.object(embeddings_local.settings, "EMBEDDINGS_THREADS", 2, create=True):
EmbeddingsWrapper(MPNET.name)
assert text_embedding.call_args.kwargs["threads"] == 2
def test_threads_omitted_when_unset(self, fake_fastembed):
text_embedding, _ = fake_fastembed
with patch.object(embeddings_local.settings, "EMBEDDINGS_THREADS", None, create=True):
EmbeddingsWrapper(MPNET.name)
assert "threads" not in text_embedding.call_args.kwargs
def test_cache_dir_forwarded_when_configured(self, fake_fastembed):
text_embedding, _ = fake_fastembed
with patch.object(embeddings_local.settings, "EMBEDDINGS_CACHE_DIR", "/models", create=True):
EmbeddingsWrapper(MPNET.name)
assert text_embedding.call_args.kwargs["cache_dir"] == "/models"
class TestEmbedding:
def test_embed_documents_returns_plain_lists(self, fake_fastembed):
_, instance = fake_fastembed
instance.embed.return_value = iter([np.array([1.0, 2.0]), np.array([3.0, 4.0])])
wrapper = EmbeddingsWrapper(MPNET.name)
assert wrapper.embed_documents(["a", "b"]) == [[1.0, 2.0], [3.0, 4.0]]
def test_embed_documents_short_circuits_on_empty_input(self, fake_fastembed):
_, instance = fake_fastembed
wrapper = EmbeddingsWrapper(MPNET.name)
instance.embed.reset_mock()
assert wrapper.embed_documents([]) == []
instance.embed.assert_not_called()
def test_embed_query_returns_a_single_vector(self, fake_fastembed):
_, instance = fake_fastembed
instance.embed.return_value = iter([np.array([0.5, 0.6])])
wrapper = EmbeddingsWrapper(MPNET.name)
assert wrapper.embed_query("hello") == [0.5, 0.6]
def test_call_dispatches_on_input_type(self, fake_fastembed):
_, instance = fake_fastembed
wrapper = EmbeddingsWrapper(MPNET.name)
instance.embed.return_value = iter([np.array([1.0])])
assert wrapper("text") == [1.0]
instance.embed.return_value = iter([np.array([1.0]), np.array([2.0])])
assert wrapper(["a", "b"]) == [[1.0], [2.0]]
def test_call_rejects_other_types(self, fake_fastembed):
wrapper = EmbeddingsWrapper(MPNET.name)
with pytest.raises(ValueError):
wrapper(42)
class TestRegistrationIsIdempotent:
def test_model_registered_once_per_process(self, fake_fastembed):
text_embedding, _ = fake_fastembed
EmbeddingsWrapper(MPNET.name)
EmbeddingsWrapper(MPNET.name)
assert text_embedding.add_custom_model.call_count == 1
class TestLengthSortedBatching:
"""Grouping by length is a throughput/memory win, but order is a contract."""
def _wrapper(self, fake_fastembed, batch_size):
_, instance = fake_fastembed
wrapper = EmbeddingsWrapper(MPNET.name)
instance.embed.side_effect = lambda texts, batch_size=None: iter(
[np.array([float(len(t))]) for t in texts]
)
mock_st_cls.return_value = mock_model
return wrapper, instance
from application.vectorstore.embeddings_local import EmbeddingsWrapper
def test_output_order_matches_input_order(self, fake_fastembed):
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)
# Each stub vector encodes its own text length, so a reordered result
# is immediately visible.
assert out == [[4.0], [1.0], [3.0], [2.0], [5.0]]
wrapper = EmbeddingsWrapper("model")
result = wrapper.embed_documents(["doc1", "doc2"])
def test_inputs_are_grouped_by_length_before_batching(self, fake_fastembed):
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]
mock_model.encode.assert_called_with(["doc1", "doc2"])
assert result == [[0.1, 0.2], [0.3, 0.4]]
def test_single_batch_is_not_reordered(self, fake_fastembed):
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)
assert instance.embed.call_args.args[0] == texts
assert out == [[4.0], [1.0], [3.0]]
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_call_with_string(self, mock_st_cls):
mock_model = MagicMock()
mock_model._first_module.return_value = MagicMock()
mock_model.get_sentence_embedding_dimension.return_value = 3
mock_model.encode.return_value = MagicMock(tolist=Mock(return_value=[0.1]))
mock_st_cls.return_value = mock_model
def test_duplicate_texts_are_handled(self, fake_fastembed):
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]]
from application.vectorstore.embeddings_local import EmbeddingsWrapper
wrapper = EmbeddingsWrapper("model")
result = wrapper("hello")
assert result == [0.1]
class TestTokenizerPadding:
"""A fixed padding width in ``tokenizer.json`` makes mixed batches ragged.
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_call_with_list(self, mock_st_cls):
mock_model = MagicMock()
mock_model._first_module.return_value = MagicMock()
mock_model.get_sentence_embedding_dimension.return_value = 3
mock_model.encode.return_value = MagicMock(
tolist=Mock(return_value=[[0.1], [0.2]])
FastEmbed calls ``enable_padding`` only when the tokenizer declares none,
so mpnet's fixed ``length: 128`` survives loading. Any batch mixing an
input longer than 128 tokens with a shorter one then produces rows of
different widths and ONNX rejects the tensor.
"""
def _tokenizer(self, padding):
tokenizer = MagicMock()
tokenizer.padding = padding
return tokenizer
def test_fixed_width_padding_is_reset_to_batch_longest(self, fake_fastembed):
_, instance = fake_fastembed
tokenizer = self._tokenizer(
{
"length": 128,
"pad_id": 1,
"pad_token": "<pad>",
"pad_type_id": 0,
"direction": "right",
"pad_to_multiple_of": None,
}
)
mock_st_cls.return_value = mock_model
instance.model.tokenizer = tokenizer
from application.vectorstore.embeddings_local import EmbeddingsWrapper
EmbeddingsWrapper(MPNET.name)
wrapper = EmbeddingsWrapper("model")
result = wrapper(["a", "b"])
assert result == [[0.1], [0.2]]
kwargs = tokenizer.enable_padding.call_args.kwargs
assert kwargs["length"] is None, "padding must follow the longest input"
# The model's own pad token must survive the reset.
assert kwargs["pad_id"] == 1
assert kwargs["pad_token"] == "<pad>"
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_call_with_invalid_type(self, mock_st_cls):
mock_model = MagicMock()
mock_model._first_module.return_value = MagicMock()
mock_model.get_sentence_embedding_dimension.return_value = 3
mock_st_cls.return_value = mock_model
def test_dynamic_padding_is_left_alone(self, fake_fastembed):
_, instance = fake_fastembed
tokenizer = self._tokenizer({"length": None, "pad_id": 0, "pad_token": "<pad>"})
instance.model.tokenizer = tokenizer
from application.vectorstore.embeddings_local import EmbeddingsWrapper
EmbeddingsWrapper(GRANITE_97M.name)
wrapper = EmbeddingsWrapper("model")
with pytest.raises(ValueError, match="Input must be a string or a list"):
wrapper(123)
tokenizer.enable_padding.assert_not_called()
@patch("application.vectorstore.embeddings_local.SentenceTransformer")
def test_trust_remote_code_default(self, mock_st_cls):
mock_model = MagicMock()
mock_model._first_module.return_value = MagicMock()
mock_model.get_sentence_embedding_dimension.return_value = 768
mock_st_cls.return_value = mock_model
def test_tokenizer_that_cannot_be_reached_is_not_fatal(self, fake_fastembed):
_, instance = fake_fastembed
instance.model = None
EmbeddingsWrapper(GRANITE_97M.name)
from application.vectorstore.embeddings_local import EmbeddingsWrapper
EmbeddingsWrapper("model")
def _repo_json(pooling_file, modules_file):
"""Stub ``_read_repo_json`` returning canned repository metadata."""
call_kwargs = mock_st_cls.call_args[1]
assert call_kwargs["trust_remote_code"] is True
def read(repo, filename):
return pooling_file if filename == embeddings_local._POOLING_CONFIG else modules_file
return read
class TestPoolingReadFromTheRepository:
"""A model's pooling is a fact its repository states, not a default.
Assuming mean pooling for a CLS model returns vectors at cosine ~0.95 to
the correct ones: no error, no dimension mismatch, just quietly worse
retrieval. These cover the shapes seen on the hub.
"""
def test_cls_pooling_is_read_rather_than_assumed(self):
with patch.object(
embeddings_local,
"_read_repo_json",
_repo_json(
{"pooling_mode_cls_token": True, "word_embedding_dimension": 384},
[{"type": "sentence_transformers.models.Transformer"},
{"type": "sentence_transformers.models.Pooling"},
{"type": "sentence_transformers.models.Normalize"}],
),
):
spec = embeddings_local._spec_for("BAAI/bge-small-en-v1.5")
assert spec.pooling == "cls"
assert spec.normalize is True
# Declared width, so no probe run is needed to learn it.
assert spec.dimension == 384
def test_missing_normalize_module_means_unnormalised(self):
"""multi-qa-mpnet-base-dot-v1 is trained on unnormalised vectors."""
with patch.object(
embeddings_local,
"_read_repo_json",
_repo_json(
{"pooling_mode_cls_token": True, "word_embedding_dimension": 768},
[{"type": "sentence_transformers.models.Transformer"},
{"type": "sentence_transformers.models.Pooling"}],
),
):
spec = embeddings_local._spec_for("sentence-transformers/multi-qa-mpnet-base-dot-v1")
assert spec.pooling == "cls"
assert spec.normalize is False
def test_dense_projection_head_is_refused(self):
"""FastEmbed would skip the projection and emit the wrong vectors."""
with patch.object(
embeddings_local,
"_read_repo_json",
_repo_json(
{"pooling_mode_cls_token": True, "word_embedding_dimension": 768},
[{"type": "sentence_transformers.models.Transformer"},
{"type": "sentence_transformers.models.Pooling"},
{"type": "sentence_transformers.models.Dense"},
{"type": "sentence_transformers.models.Normalize"}],
),
):
with pytest.raises(RuntimeError) as excinfo:
embeddings_local._spec_for("sentence-transformers/LaBSE")
message = str(excinfo.value)
assert "LaBSE" in message
assert "Dense" in message
def test_unsupported_pooling_mode_falls_back_rather_than_lying(self):
with patch.object(
embeddings_local,
"_read_repo_json",
_repo_json({"pooling_mode_max_tokens": True}, []),
):
spec = embeddings_local._spec_for("some-org/max-pooled")
assert spec.pooling == embeddings_local._FALLBACK_POOLING
assert spec.dimension == 0
def test_repository_without_metadata_keeps_the_assumption(self):
spec = embeddings_local._spec_for("some-org/plain-onnx-export")
assert spec.pooling == embeddings_local._FALLBACK_POOLING
assert spec.normalize is True
assert spec.dimension == 0
def test_registry_wins_over_the_repository(self):
"""A described model is never re-read; the registry is the answer."""
read = MagicMock()
with patch.object(embeddings_local, "_read_repo_json", read):
spec = embeddings_local._spec_for(MPNET.name)
assert spec is MPNET
read.assert_not_called()
class TestPoolingOverrides:
def test_settings_override_what_the_repository_declares(self):
with patch.object(
embeddings_local,
"_read_repo_json",
_repo_json(
{"pooling_mode_mean_tokens": True, "word_embedding_dimension": 768},
[{"type": "sentence_transformers.models.Normalize"}],
),
):
with patch.object(embeddings_local.settings, "EMBEDDINGS_POOLING", "cls"), \
patch.object(embeddings_local.settings, "EMBEDDINGS_NORMALIZE", False):
spec = embeddings_local._spec_for("some-org/mislabelled")
assert spec.pooling == "cls"
assert spec.normalize is False
def test_a_meaningless_override_is_ignored(self):
with patch.object(embeddings_local.settings, "EMBEDDINGS_POOLING", "banana"):
spec = embeddings_local._spec_for("some-org/plain-onnx-export")
assert spec.pooling == embeddings_local._FALLBACK_POOLING
+68 -7
View File
@@ -5,11 +5,13 @@ than asserting that calls were forwarded to a mock.
"""
import json
from unittest.mock import Mock, patch
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock, patch
import pytest
from application.storage.local import LocalStorage
from application.vectorstore.faiss import FaissStore
class _FakeEmbeddings:
@@ -185,15 +187,20 @@ class TestFaissStoreAssertEmbeddingDimensions:
with pytest.raises(ValueError, match="Embedding dimension mismatch"):
populated.assert_embedding_dimensions(Mock(dimension=768))
def test_missing_dimension_attr_raises(self, populated):
def test_unknown_dimension_defers_rather_than_raising(self, populated):
"""A remote model reports no width until its first call.
Refusing to open the index in that window would break startup for a
perfectly valid remote configuration, so an unknown width is deferred,
not treated as a mismatch.
"""
with patch("application.vectorstore.faiss.settings") as mock_settings:
mock_settings.EMBEDDINGS_NAME = (
"huggingface_sentence-transformers/all-mpnet-base-v2"
)
embeddings = Mock()
del embeddings.dimension
with pytest.raises(AttributeError, match="'dimension' attribute not found"):
populated.assert_embedding_dimensions(embeddings)
embeddings.dimension = None
assert populated.assert_embedding_dimensions(embeddings) is None
def test_dimension_match_passes(self, populated):
with patch("application.vectorstore.faiss.settings") as mock_settings:
@@ -202,10 +209,22 @@ class TestFaissStoreAssertEmbeddingDimensions:
)
assert populated.assert_embedding_dimensions(Mock(dimension=3)) is None
def test_non_huggingface_skips_dimension_check(self, populated):
def test_mismatch_is_caught_for_every_model_not_just_mpnet(self, populated):
"""The check used to run only when EMBEDDINGS_NAME was mpnet.
That skipped exactly the case it exists for: an index built with one
model being opened under a different one.
"""
with patch("application.vectorstore.faiss.settings") as mock_settings:
mock_settings.EMBEDDINGS_NAME = "openai_text-embedding-ada-002"
assert populated.assert_embedding_dimensions(Mock(dimension=1536)) is None
with pytest.raises(ValueError, match="Embedding dimension mismatch"):
populated.assert_embedding_dimensions(Mock(dimension=1536))
def test_mismatch_message_points_at_the_reembed_script(self, populated):
with patch("application.vectorstore.faiss.settings") as mock_settings:
mock_settings.EMBEDDINGS_NAME = "granite-311m"
with pytest.raises(ValueError, match="application.scripts.reembed"):
populated.assert_embedding_dimensions(Mock(dimension=768))
@pytest.mark.unit
@@ -226,3 +245,45 @@ class TestGetVectorstore:
with pytest.raises(ValueError, match="Invalid source_id path"):
get_vectorstore(bad)
class TestBuildFromDocumentsBatching:
"""The full rebuild path: honour caller ids and embed in batches."""
def _docs(self, n):
return [
SimpleNamespace(page_content=f"chunk {i}", metadata={"i": i})
for i in range(n)
]
def _store(self, embed):
store = FaissStore.__new__(FaissStore)
store.embeddings = MagicMock()
store.embeddings.embed_documents.side_effect = embed
store.documents = {}
store.index_to_docstore_id = {}
store.index = None
return store
def test_supplied_ids_are_used(self):
store = self._store(lambda texts: [[float(len(texts))] * 2 for _ in texts])
store._build_from_documents(self._docs(3), ids=["a", "b", "c"])
assert list(store.documents) == ["a", "b", "c"]
def test_embedding_is_split_into_batches(self):
sizes = []
def embed(texts):
sizes.append(len(texts))
return [[1.0, 2.0] for _ in texts]
store = self._store(embed)
store._build_from_documents(self._docs(5), batch_size=2)
assert sizes == [2, 2, 1]
assert len(store.index_to_docstore_id) == 5
def test_defaults_are_unchanged(self):
store = self._store(lambda texts: [[1.0, 2.0] for _ in texts])
store._build_from_documents(self._docs(3))
assert len(store.documents) == 3
assert all(len(k) == 36 for k in store.documents), "uuid4 ids by default"
+117
View File
@@ -0,0 +1,117 @@
"""The registry is the single source of truth for embedding-model facts."""
import os
import pytest
from application.vectorstore import model_registry as reg
class TestResolve:
def test_resolves_canonical_name(self):
assert reg.resolve(reg.MPNET.name) is reg.MPNET
@pytest.mark.parametrize(
"alias",
[
"huggingface_sentence-transformers-all-mpnet-base-v2",
"sentence-transformers/all-mpnet-base-v2",
"all-mpnet-base-v2",
],
)
def test_resolves_legacy_spellings_of_mpnet(self, alias):
"""Every spelling the old factory dict accepted must still work."""
assert reg.resolve(alias) is reg.MPNET
def test_resolution_is_case_insensitive_and_trims(self):
assert reg.resolve(" GRANITE-311M ") is reg.GRANITE_311M
def test_unknown_name_is_none_not_an_error(self):
"""Unknown names are a valid configuration: an arbitrary HF repo."""
assert reg.resolve("some-org/some-model") is None
def test_empty_and_none_resolve_to_none(self):
assert reg.resolve(None) is None
assert reg.resolve("") is None
class TestModelFacts:
def test_granite_311m_matches_the_existing_column_width(self):
"""768 is what makes granite a drop-in for an mpnet index."""
assert reg.GRANITE_311M.dimension == reg.MPNET.dimension == 768
def test_granite_97m_is_narrower(self):
assert reg.GRANITE_97M.dimension == 384
def test_pooling_and_normalisation_are_recorded(self):
assert reg.MPNET.pooling == "mean"
assert reg.GRANITE_311M.pooling == "cls"
assert all(m.normalize for m in reg.MODELS)
def test_granite_context_is_much_wider_than_mpnet(self):
assert reg.GRANITE_311M.max_input_tokens == 32768
assert reg.MPNET.max_input_tokens == 384
def test_local_runners_carry_a_repo_and_onnx_file(self):
for model in reg.MODELS:
if model.provider == "fastembed":
assert model.repo, f"{model.name} has no repo"
assert model.onnx_file, f"{model.name} has no onnx_file"
def test_openai_model_needs_no_local_artifacts(self):
assert reg.OPENAI_ADA_002.provider == "openai"
assert reg.OPENAI_ADA_002.repo is None
class TestHelpers:
def test_dimension_for_known_and_unknown(self):
assert reg.dimension_for(reg.GRANITE_97M.name) == 384
assert reg.dimension_for("nope/nope") is None
def test_max_input_tokens_for_known_and_unknown(self):
assert reg.max_input_tokens_for("granite-311m") == 32768
assert reg.max_input_tokens_for("nope/nope") is None
def test_known_names_lists_canonical_spellings(self):
names = reg.known_names()
assert reg.MPNET.name in names
assert reg.GRANITE_311M.name in names
def test_defaults_point_at_the_intended_models(self):
"""Existing installs stay on mpnet; new installs get granite."""
assert reg.DEFAULT_LEGACY == reg.MPNET.name
assert reg.DEFAULT_NEW_INSTALL == reg.GRANITE_311M.name
def test_no_alias_collisions_between_models(self):
seen = {}
for model in reg.MODELS:
for key in (model.name, *model.aliases):
assert key.lower() not in seen, f"{key} claimed twice"
seen[key.lower()] = model
class TestRegistryMatchesTheHub:
"""Every registry entry restates facts the model's repository already holds.
Restating them is what makes an offline install work, but a restatement can
drift from its source and nothing else would notice: wrong pooling is not a
crash, only worse retrieval. Opt in with ``DOCSGPT_HUB_TESTS=1``; needs
network.
"""
@pytest.mark.skipif(
os.environ.get("DOCSGPT_HUB_TESTS") != "1",
reason="set DOCSGPT_HUB_TESTS=1 to check the registry against the hub",
)
@pytest.mark.parametrize(
"model", [m for m in reg.MODELS if m.provider == "fastembed"]
)
def test_entry_matches_repository_metadata(self, model):
from application.vectorstore.embeddings_local import _describe_from_repo
described = _describe_from_repo(model.repo)
assert described is not None, f"{model.repo} declares no pooling metadata"
assert described.pooling == model.pooling
assert described.normalize == model.normalize
if described.dimension:
assert described.dimension == model.dimension
@@ -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()
@@ -83,3 +83,115 @@ def test_query_path_is_truncated(monkeypatch):
sent = captured["payload"]["input"]
assert sent == enc.decode(enc.encode(long_text)[:10])
class TestInputLimitResolution:
"""The cap falls back to the model's own context window."""
def _remote(self, model_name):
from application.vectorstore.base import RemoteEmbeddings
return RemoteEmbeddings(
api_url="http://embeddings", model_name=model_name, api_key=None
)
def test_explicit_setting_wins(self, monkeypatch):
from application.vectorstore import base
monkeypatch.setattr(base.settings, "EMBEDDINGS_MAX_INPUT_TOKENS", 123)
assert self._remote("granite-311m")._resolve_input_limit() == 123
def test_registered_model_supplies_its_own_ceiling(self, monkeypatch):
from application.vectorstore import base
monkeypatch.setattr(base.settings, "EMBEDDINGS_MAX_INPUT_TOKENS", None)
monkeypatch.setattr(base, "_embeddings_name_is_explicit", lambda: True)
assert self._remote("granite-311m")._resolve_input_limit() == 32768
# mpnet genuinely stops at 384; sending more is paid for and discarded.
assert self._remote("all-mpnet-base-v2")._resolve_input_limit() == 384
def test_default_model_name_lends_the_server_no_ceiling(self, monkeypatch):
"""An unset EMBEDDINGS_NAME describes nothing about the remote server.
The name is only forwarded as the ``model`` field. Letting the
settings default contribute mpnet's 384-token window would clip every
chunk on a server that may well serve a 32k-context model.
"""
from application.vectorstore import base
monkeypatch.setattr(base.settings, "EMBEDDINGS_MAX_INPUT_TOKENS", None)
monkeypatch.setattr(base, "_embeddings_name_is_explicit", lambda: False)
assert self._remote("all-mpnet-base-v2")._resolve_input_limit() is None
def test_explicit_setting_still_wins_over_an_unset_name(self, monkeypatch):
from application.vectorstore import base
monkeypatch.setattr(base.settings, "EMBEDDINGS_MAX_INPUT_TOKENS", 512)
monkeypatch.setattr(base, "_embeddings_name_is_explicit", lambda: False)
assert self._remote("all-mpnet-base-v2")._resolve_input_limit() == 512
def test_unknown_model_stays_unlimited(self, monkeypatch):
from application.vectorstore import base
monkeypatch.setattr(base.settings, "EMBEDDINGS_MAX_INPUT_TOKENS", None)
assert self._remote("some-org/mystery")._resolve_input_limit() is None
def test_non_positive_setting_falls_through_to_the_registry(self, monkeypatch):
from application.vectorstore import base
monkeypatch.setattr(base.settings, "EMBEDDINGS_MAX_INPUT_TOKENS", 0)
monkeypatch.setattr(base, "_embeddings_name_is_explicit", lambda: True)
assert self._remote("granite-97m")._resolve_input_limit() == 32768
def test_dimension_is_taken_from_the_registry(self):
assert self._remote("granite-97m").dimension == 384
assert self._remote("granite-311m").dimension == 768
def test_unknown_model_dimension_is_probed_not_assumed(self):
"""The old hardcoded 768 made the probe below unreachable."""
assert self._remote("some-org/mystery").dimension is None
class TestEmbeddingsNameIsExplicit:
"""Which names count as chosen decides whether a remote server inherits a
context window it may not have.
The tests above monkeypatch the predicate, so they cannot see it being
wrong. These drive the real one.
"""
def _default(self):
from application.core.settings import Settings
return Settings.model_fields["EMBEDDINGS_NAME"].default
def test_the_default_name_is_not_a_choice(self, monkeypatch):
"""Every setup script has always written this value unconditionally.
Reading it as deliberate lends the server mpnet's 384-token window and
clips ~80% off every chunk, silently, on upgrade. ``model_fields_set``
could not tell the difference: pydantic marks a field set for anything
that reached it, ``.env`` included.
"""
monkeypatch.setattr(base.settings, "EMBEDDINGS_NAME", self._default())
assert base._embeddings_name_is_explicit() is False
def test_a_different_name_is_a_choice(self, monkeypatch):
monkeypatch.setattr(
base.settings,
"EMBEDDINGS_NAME",
"ibm-granite/granite-embedding-311m-multilingual-r2",
)
assert base._embeddings_name_is_explicit() is True
def test_dotenv_written_default_does_not_clip(self, monkeypatch):
"""End to end: the common upgrade path must send the full chunk."""
monkeypatch.setattr(base.settings, "EMBEDDINGS_MAX_INPUT_TOKENS", None)
monkeypatch.setattr(base.settings, "EMBEDDINGS_NAME", self._default())
captured = _capture_post(monkeypatch)
long_text = " ".join(["word"] * 1000)
emb = RemoteEmbeddings(api_url="http://embeddings", model_name=self._default())
emb.embed_documents([long_text])
assert captured["payload"]["input"][0] == long_text
+130
View File
@@ -326,3 +326,133 @@ class TestAttachmentZipBombGuard:
# BadZipFile → return quietly; the format parser surfaces a clean error.
worker._reject_attachment_zip_bomb(str(path))
@pytest.mark.unit
class TestAttachmentTypeGuard:
"""A binary with no parser must fail rather than be read as text.
``SimpleDirectoryReader`` falls through to a plain-text ``open()`` for a
suffix it has no parser for, which is how a phone-uploaded ``.mp4`` was
once stored as 100k tokens of binary with ``extraction.status == "ok"``.
The same fallthrough is what makes a ``.py`` or ``.log`` attachment work,
so the guard admits those on content.
"""
def test_binary_without_a_parser_fails_instead_of_being_read_as_text(
self, pg_conn, patch_worker_db, task_self, monkeypatch, tmp_path
):
from application import worker
local_path = tmp_path / "clip.mp4"
local_path.write_bytes(b"\x00\x00\x00\x18ftypisom\x00\x00\x02\x00isomavc1")
events = []
fake_storage = MagicMock(name="storage")
fake_storage.process_file.side_effect = lambda path, callback: callback(
str(local_path)
)
monkeypatch.setattr(worker.StorageCreator, "get_storage", lambda: fake_storage)
monkeypatch.setattr(
worker,
"get_default_file_extractor",
lambda ocr_enabled=False, pdf_text_fast_path=False: {},
)
monkeypatch.setattr(
worker,
"publish_user_event",
lambda user, name, payload, **kwargs: events.append((name, payload)),
)
file_info = {
"filename": "clip.mp4",
"attachment_id": "507f1f77bcf86cd799439012",
"path": "uploads/user1/attachments/clip.mp4",
"metadata": {"source": "chat"},
}
with pytest.raises(
worker.AttachmentRejectedError, match=r"Unsupported file type: \.mp4"
):
worker.attachment_worker(task_self, file_info, "user1")
failed = [payload for name, payload in events if name == "attachment.failed"]
assert failed and failed[0]["error"] == "Unsupported file type: .mp4"
def test_suffix_the_loaded_extractor_cannot_parse_is_refused(
self, pg_conn, patch_worker_db, task_self, monkeypatch, tmp_path
):
"""The guard judges against the parser table actually loaded.
Without docling the fallback extractor has no .webp parser, so a
.webp the route admitted on its name would otherwise be opened as
plain text here.
"""
from application import worker
local_path = tmp_path / "scan.webp"
local_path.write_bytes(b"RIFF\x00\x00\x00\x00WEBPVP8 " + bytes(range(256)))
events = []
fake_storage = MagicMock(name="storage")
fake_storage.process_file.side_effect = lambda path, callback: callback(
str(local_path)
)
monkeypatch.setattr(worker.StorageCreator, "get_storage", lambda: fake_storage)
# The docling-less fallback table: images are handled, .webp is not.
monkeypatch.setattr(
worker,
"get_default_file_extractor",
lambda ocr_enabled=False, pdf_text_fast_path=False: {".png": object()},
)
monkeypatch.setattr(
worker,
"publish_user_event",
lambda user, name, payload, **kwargs: events.append((name, payload)),
)
file_info = {
"filename": "scan.webp",
"attachment_id": "507f1f77bcf86cd799439014",
"path": "uploads/user1/attachments/scan.webp",
"metadata": {"source": "chat"},
}
with pytest.raises(
worker.AttachmentRejectedError, match=r"Unsupported file type: \.webp"
):
worker.attachment_worker(task_self, file_info, "user1")
failed = [payload for name, payload in events if name == "attachment.failed"]
assert failed and failed[0]["error"] == "Unsupported file type: .webp"
def test_text_without_a_parser_is_parsed(
self, pg_conn, patch_worker_db, task_self, monkeypatch, tmp_path
):
from application import worker
local_path = tmp_path / "server.log"
local_path.write_text("2026-09-02 ERROR boom\n", encoding="utf-8")
fake_storage = MagicMock(name="storage")
fake_storage.process_file.side_effect = lambda path, callback: callback(
str(local_path)
)
monkeypatch.setattr(worker.StorageCreator, "get_storage", lambda: fake_storage)
monkeypatch.setattr(
worker,
"get_default_file_extractor",
lambda ocr_enabled=False, pdf_text_fast_path=False: {},
)
file_info = {
"filename": "server.log",
"attachment_id": "507f1f77bcf86cd799439013",
"path": "uploads/user1/attachments/server.log",
"metadata": {"source": "chat"},
}
result = worker.attachment_worker(task_self, file_info, "user1")
assert result["filename"] == "server.log"
assert result["token_count"] > 0
+40
View File
@@ -100,6 +100,46 @@ class TestReembedWikiPageWorker:
)
assert result == {"status": "embedded", "added": 2, "deleted": 3}
def test_a_reembed_stamps_the_model_on_the_source(
self, pg_conn, patch_worker_db, task_self, monkeypatch
):
"""Wiki sources are created with ``model`` NULL, which the boot mismatch
check reads as the legacy model and reports as stale on every startup.
Stamping here heals a source created before it was recorded.
"""
from application import worker
from application.core.settings import settings
source_id = _seed_source(pg_conn)
assert SourcesRepository(pg_conn).get_any(source_id, "alice")["model"] is None
_patch_store(monkeypatch, MagicMock(name="vector_store"))
_patch_repo(monkeypatch, {"content": "body", "title": "T"})
_patch_chunker(monkeypatch, [Document(text="c1")])
worker.reembed_wiki_page_worker(
task_self, source_id, "p.md", "hash-1", "alice"
)
row = SourcesRepository(pg_conn).get_any(source_id, "alice")
assert row["model"] == settings.EMBEDDINGS_NAME
def test_a_purge_leaves_the_model_alone(
self, pg_conn, patch_worker_db, task_self, monkeypatch
):
"""Deleting a page embeds nothing, so it claims nothing about the model."""
from application import worker
source_id = _seed_source(pg_conn)
_patch_store(monkeypatch, MagicMock(name="vector_store"))
_patch_repo(monkeypatch, None)
worker.reembed_wiki_page_worker(
task_self, source_id, "gone.md", "hash-x", "alice"
)
assert SourcesRepository(pg_conn).get_any(source_id, "alice")["model"] is None
def test_page_missing_purges(
self, pg_conn, patch_worker_db, task_self, monkeypatch
):