mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-04 06:12:49 +00:00
Review follow-up. The per-group secret validators normalised a hand-picked list of API keys, which left other optional credentials and overrides (OPEN_ROUTER_API_KEY, S3 and Daytona keys, ELASTIC_PASSWORD, the OIDC trio, connector client ids, MICROSOFT_AUTHORITY, MCP_OAUTH_REDIRECT_URI) holding the literal "None" or "" a .env file spells "unset" with, so truthiness checks and fallbacks downstream saw a value. One rule on the group base replaces those lists: every Optional[str] field maps "", "None" and whitespace to None and strips real values. Plain str fields are left alone. The OIDC required-settings check therefore also rejects those spellings. EMBEDDINGS_POOLING is Literal["cls", "mean"] with case-insensitive parsing; its consumer silently ignored anything else. Bounds added where the consumer rejects or misbehaves on the value: SCHEDULE_RUN_OUTPUT_RETENTION_DAYS and MESSAGE_EVENTS_RETENTION_DAYS (the cleanup repositories raise on <= 0), EMBEDDINGS_DELEGATE_TIMEOUT, the remote-device idle/pairing/invocation TTLs and CELERY_VISIBILITY_TIMEOUT (> 0), REMOTE_DEVICE_CMD_QUEUE_TTL_SECONDS (> 605, the documented drain deadline), GRAPHRAG_MAX_CHUNKS_FOR_EXTRACTION (>= 0; negative would slice the pending list from the end). The generated reference now renders generic type arguments (dict[str, int] rather than dict).
89 lines
3.7 KiB
Python
89 lines
3.7 KiB
Python
"""Embedding model selection and where it runs."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import Literal, Optional
|
|
|
|
from pydantic import Field, field_validator
|
|
|
|
from docsgpt.core.paths import home_dir
|
|
from docsgpt.core.settings._shared import SettingsGroup, normalize_choice
|
|
|
|
|
|
class EmbeddingsSettings(SettingsGroup):
|
|
"""The embedding model, remote or local, and the batching around it."""
|
|
|
|
EMBEDDINGS_NAME: str = Field(
|
|
default="huggingface_sentence-transformers/all-mpnet-base-v2",
|
|
description=(
|
|
"Embedding model. The legacy model is the default 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 "
|
|
"docsgpt.scripts.reembed."
|
|
),
|
|
)
|
|
EMBEDDINGS_BASE_URL: Optional[str] = Field(
|
|
default=None, description="Remote embeddings API URL (OpenAI-compatible)."
|
|
)
|
|
EMBEDDINGS_KEY: Optional[str] = Field(
|
|
default=None, description="API key for embeddings (with OpenAI, the same value as API_KEY)."
|
|
)
|
|
EMBEDDINGS_MAX_INPUT_TOKENS: Optional[int] = Field(
|
|
default=None, description="Truncate each remote embed input to N tokens (overflow is lost)."
|
|
)
|
|
EMBEDDINGS_BATCH_SIZE: int = Field(
|
|
default=32, ge=1, description="Chunks per store transaction and per remote embed request."
|
|
)
|
|
EMBEDDINGS_MODEL_BATCH_SIZE: int = Field(
|
|
default=1,
|
|
ge=1,
|
|
description=(
|
|
"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_THREADS: Optional[int] = Field(
|
|
default=None,
|
|
description=(
|
|
"Intra-op threads for the local ONNX runner; unset uses every core. It scales sub-linearly, so "
|
|
"several single-threaded workers beat one many-threaded process on the same cores."
|
|
),
|
|
)
|
|
EMBEDDINGS_CACHE_DIR: Optional[str] = Field(
|
|
default_factory=lambda: str(home_dir() / "models"),
|
|
description=(
|
|
"Where embedding models and their tokenizers are cached. Persistent by default: FastEmbed's own "
|
|
"default is the temp dir."
|
|
),
|
|
)
|
|
EMBEDDINGS_POOLING: Optional[Literal["cls", "mean"]] = Field(
|
|
default=None,
|
|
description=(
|
|
'Pooling strategy ("cls" or "mean"). Read from the model\'s own repository; set only for a '
|
|
"repository that declares none, or to override what it declares."
|
|
),
|
|
)
|
|
EMBEDDINGS_NORMALIZE: Optional[bool] = Field(
|
|
default=None,
|
|
description=(
|
|
"L2-normalise embeddings. Read from the model's own repository; set only for a repository that "
|
|
"declares nothing, or to override what it declares."
|
|
),
|
|
)
|
|
EMBEDDINGS_DELEGATE_TO_WORKER: bool = Field(
|
|
default=True,
|
|
description=(
|
|
"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_QUEUE: str = Field(default="embeddings", description="Celery queue the embed task is routed to.")
|
|
EMBEDDINGS_DELEGATE_TIMEOUT: int = Field(
|
|
default=60, gt=0, description="Seconds the API waits for the worker to return an embedding."
|
|
)
|
|
|
|
@field_validator("EMBEDDINGS_POOLING", mode="before")
|
|
@classmethod
|
|
def _normalize_pooling(cls, v):
|
|
return normalize_choice(v)
|