diff --git a/README.md b/README.md index a2141ce..a6f819a 100644 --- a/README.md +++ b/README.md @@ -94,7 +94,7 @@ A high-performance, multi-repo Model Context Protocol (MCP) server providing **s - **Local Storage**: Managed file explorer, direct file upload modal with folder categorization, and file preview/replacement. - **Ingestion Catalog**: Unified multi-source explorer with source type filters, repository lookup, and file listings. - **Search & Inspector**: Interactive live hybrid search tester with RRF score previews, target type toggle (Code vs Docs), and syntax highlighted results. - - **Settings**: Vector Database manager (pgvector, Qdrant, & ChromaDB switcher & connection tester), multi-provider token cards, GitHub rate limit monitor, and interactive Custom Git Host Credential Vault table/modal. + - **Settings**: Vector Database manager (pgvector, Qdrant, & ChromaDB switcher & connection tester), LiteLLM Model Discovery with dynamic categorized model dropdowns (Embeddings, Vision OCR, and Chat models), multi-provider token cards, GitHub rate limit monitor, and interactive Custom Git Host Credential Vault table/modal. - **Diagnostics & Logs**: Real-time log viewer with level filtering (ALL, INFO, WARNING, ERROR, DEBUG), keyword search, traceback modal/drawer, and buffer clearing. --- diff --git a/REQUIREMENTS.md b/REQUIREMENTS.md index 58c8518..ca22170 100644 --- a/REQUIREMENTS.md +++ b/REQUIREMENTS.md @@ -2,7 +2,7 @@ > **Note:** This document is automatically generated and verified against the live test suite by `scripts/generate_requirements.py` and `tests/backend/test_requirements_sync.py`. -**Test Verification Baseline:** **905 Automated Tests** (598 Pytest Backend + 261 Vitest Frontend + 46 Playwright E2E). +**Test Verification Baseline:** **923 Automated Tests** (611 Pytest Backend + 266 Vitest Frontend + 46 Playwright E2E). --- @@ -949,6 +949,15 @@ persisting all records and vector points correctly across multiple flushes._ - `test_incremental_pipeline_clone_error_resilience` - _Verifies that a failure during shallow clone records an error in git_repositories and leaves the prior indexed state intact without data loss._ +#### `tests/test_litellm_service.py` (7 tests) +- `test_discover_models_success` +- `test_discover_models_timeout` +- `test_discover_models_connect_error` +- `test_discover_models_http_401_unauthorized` +- `test_discover_models_http_500_error` +- `test_discover_models_url_normalization` +- `test_discover_models_default_resolution` + #### `tests/test_local_storage_indexing.py` (4 tests) - `test_incremental_indexing_on_save` - `test_incremental_indexing_code_file` @@ -980,6 +989,16 @@ and leaves the prior indexed state intact without data loss._ - `test_what_is_ingested_detailed_with_data` - `test_tool_registration` +#### `tests/test_model_discovery_api.py` (2 tests) +- `test_api_discover_models_endpoint` +- `test_api_embedding_settings_get_and_post_with_models` + +#### `tests/test_model_metadata_persistence.py` (4 tests) +- `test_embedding_db_config_defaults` +- `test_env_variable_fallbacks` +- `test_set_embedding_db_config_persists_to_system_metadata` +- `test_update_embedding_config_service` + #### `tests/test_pdf_extractor.py` (8 tests) - `test_extract_digital_pdf_text` - `test_extract_scanned_pdf_triggers_vision_ocr` @@ -1138,11 +1157,16 @@ and leaves the prior indexed state intact without data loss._ - displays error toast when log fetching fails - renders responsive layout elements for toolbar, search input, and log entry stream -#### `EmbeddingSettings.test.tsx` (4 tests) +#### `EmbeddingSettings.test.tsx` (9 tests) - renders loading state when embedding configuration is not yet loaded - renders active status with hardware metrics and local model parameters - handles provider switch to API and updates form fields - handles changes to CPU threads and batch size +- renders model discovery controls and triggers onDiscoverModels when button clicked +- renders discovered model dropdowns and allows selecting models +- displays discovery error banner when LiteLLM is unreachable +- switches to manual input when Custom is selected from dropdown or link clicked +- disables discover button and shows spinner while isDiscovering is true #### `GitRepoManager.test.tsx` (13 tests) - renders repository list with status badges, auto-sync buttons, and details diff --git a/app/api/routers/settings.py b/app/api/routers/settings.py index b1cd022..d14d265 100644 --- a/app/api/routers/settings.py +++ b/app/api/routers/settings.py @@ -18,6 +18,7 @@ import app.services.vector_store as vs_service import app.services.indexing as idx_service import app.services.embeddings as emb_service +import app.services.litellm_service as litellm_service logger = logging.getLogger("contextcortex.api") @@ -116,6 +117,8 @@ def _count(table): "embedding_provider": emb_cfg["provider"], "dense_model": emb_cfg["dense_model"], "sparse_model": emb_cfg["sparse_model"], + "vision_ocr_model": emb_cfg.get("vision_ocr_model"), + "chat_model": emb_cfg.get("chat_model"), "embedding_threads": emb_cfg["threads"], "embedding_batch_size": emb_cfg["batch_size"], "system_cpus": emb_cfg.get("system_cpus", 2), @@ -319,6 +322,15 @@ def _reindex(): logger.error(f"Error switching vector store backend: {e}") return JSONResponse(status_code=500, content={"status": "error", "error": str(e), "message": str(e)}) +@router.get("/admin/api/models/discover") +async def api_discover_models(url: Optional[str] = None, api_key: Optional[str] = None): + try: + res = await litellm_service.discover_models(url=url, api_key=api_key) + return res + except Exception as e: + logger.error(f"Error discovering models: {e}") + return JSONResponse(status_code=500, content={"status": "error", "error": str(e), "message": str(e)}) + @router.get("/admin/api/settings/embedding") async def api_get_embedding_settings(): try: @@ -339,6 +351,8 @@ async def api_save_embedding_settings(payload: EmbeddingSettingsRequest): batch_size=payload.batch_size, litellm_url=payload.litellm_url, litellm_api_key=payload.litellm_api_key, + vision_ocr_model=payload.vision_ocr_model, + chat_model=payload.chat_model, ) return { "status": "success", diff --git a/app/models/schemas.py b/app/models/schemas.py index 2b4289d..1bb8e59 100644 --- a/app/models/schemas.py +++ b/app/models/schemas.py @@ -170,6 +170,8 @@ class EmbeddingSettingsRequest(BaseModel): batch_size: Optional[int] = Field(default=None, ge=1, le=1024) litellm_url: Optional[str] = None litellm_api_key: Optional[str] = None + vision_ocr_model: Optional[str] = None + chat_model: Optional[str] = None # Topology Graph Models class TopologyNode(BaseModel): diff --git a/app/services/database/__init__.py b/app/services/database/__init__.py index a30ebcc..1d19857 100644 --- a/app/services/database/__init__.py +++ b/app/services/database/__init__.py @@ -36,6 +36,8 @@ detect_system_resources, get_embedding_db_config, set_embedding_db_config, + get_vision_ocr_model, + get_chat_model, ) from app.services.database.credentials import ( list_git_host_credentials, @@ -101,6 +103,8 @@ "detect_system_resources", "get_embedding_db_config", "set_embedding_db_config", + "get_vision_ocr_model", + "get_chat_model", "list_git_host_credentials", "get_git_host_credential", "save_git_host_credential", diff --git a/app/services/database/connection.py b/app/services/database/connection.py index 01a09a7..521e4cd 100644 --- a/app/services/database/connection.py +++ b/app/services/database/connection.py @@ -224,6 +224,9 @@ def _resolve_default_embedding_config(conn: Optional[Any] = None) -> Dict[str, A litellm_url = os.getenv("LITELLM_URL", "http://litellm:4000/v1").strip() litellm_api_key = os.getenv("LITELLM_API_KEY", "dummy").strip() + vision_ocr_model = os.getenv("VISION_OCR_MODEL", "gemini-2.5-flash").strip() + chat_model = os.getenv("CHAT_MODEL", "gemini-2.5-flash").strip() + return { "provider": provider, "dense_model": dense_model, @@ -232,6 +235,8 @@ def _resolve_default_embedding_config(conn: Optional[Any] = None) -> Dict[str, A "batch_size": max(1, batch_size), "litellm_url": litellm_url, "litellm_api_key": litellm_api_key, + "vision_ocr_model": vision_ocr_model, + "chat_model": chat_model, "system_cpus": sys_res["cpus"], "system_memory_gb": sys_res["memory_gb"], } @@ -245,6 +250,8 @@ def get_embedding_db_config() -> Dict[str, Any]: batch_size_str = get_metadata("embedding_batch_size") litellm_url = get_metadata("embedding_litellm_url") litellm_api_key = get_metadata("embedding_litellm_api_key") + vision_ocr_model = get_metadata("vision_ocr_model") or get_metadata("embedding_vision_ocr_model") + chat_model = get_metadata("chat_model") or get_metadata("embedding_chat_model") default_cfg = _resolve_default_embedding_config() @@ -259,6 +266,8 @@ def get_embedding_db_config() -> Dict[str, Any]: "batch_size": max(1, batch_size), "litellm_url": (litellm_url or default_cfg["litellm_url"]).strip(), "litellm_api_key": litellm_api_key or default_cfg["litellm_api_key"], + "vision_ocr_model": (vision_ocr_model or default_cfg["vision_ocr_model"]).strip(), + "chat_model": (chat_model or default_cfg["chat_model"]).strip(), "system_cpus": default_cfg["system_cpus"], "system_memory_gb": default_cfg["system_memory_gb"], } @@ -272,6 +281,8 @@ def set_embedding_db_config( batch_size: Optional[int] = None, litellm_url: Optional[str] = None, litellm_api_key: Optional[str] = None, + vision_ocr_model: Optional[str] = None, + chat_model: Optional[str] = None, ): if provider is not None: set_metadata("embedding_provider", provider.lower().strip()) @@ -287,3 +298,25 @@ def set_embedding_db_config( set_metadata("embedding_litellm_url", litellm_url.strip()) if litellm_api_key is not None: set_metadata("embedding_litellm_api_key", litellm_api_key.strip()) + if vision_ocr_model is not None: + set_metadata("vision_ocr_model", vision_ocr_model.strip()) + set_metadata("embedding_vision_ocr_model", vision_ocr_model.strip()) + if chat_model is not None: + set_metadata("chat_model", chat_model.strip()) + set_metadata("embedding_chat_model", chat_model.strip()) + + +def get_vision_ocr_model() -> str: + """Returns stored vision OCR model or fallback to env/default.""" + stored = get_metadata("vision_ocr_model") or get_metadata("embedding_vision_ocr_model") + if stored and stored.strip(): + return stored.strip() + return (os.getenv("VISION_OCR_MODEL") or "gemini-2.5-flash").strip() + + +def get_chat_model() -> str: + """Returns stored chat model or fallback to env/default.""" + stored = get_metadata("chat_model") or get_metadata("embedding_chat_model") + if stored and stored.strip(): + return stored.strip() + return (os.getenv("CHAT_MODEL") or "gemini-2.5-flash").strip() diff --git a/app/services/embeddings.py b/app/services/embeddings.py index b9c67d2..1ead907 100644 --- a/app/services/embeddings.py +++ b/app/services/embeddings.py @@ -23,6 +23,8 @@ def _get_db_emb_config() -> Dict[str, Any]: EMBEDDING_BATCH_SIZE = _initial_cfg["batch_size"] LITELLM_URL = _initial_cfg["litellm_url"] LITELLM_API_KEY = _initial_cfg.get("litellm_api_key", os.getenv("LITELLM_API_KEY", "dummy")) +VISION_OCR_MODEL = _initial_cfg.get("vision_ocr_model", os.getenv("VISION_OCR_MODEL", "gemini-2.5-flash")) +CHAT_MODEL = _initial_cfg.get("chat_model", os.getenv("CHAT_MODEL", "gemini-2.5-flash")) _dense_model = None _sparse_model = None @@ -38,10 +40,13 @@ def init_embeddings( batch_size: Optional[int] = None, litellm_url: Optional[str] = None, litellm_api_key: Optional[str] = None, + vision_ocr_model: Optional[str] = None, + chat_model: Optional[str] = None, ): global _dense_model, _sparse_model, _openai_client global EMBEDDING_PROVIDER, DENSE_MODEL_NAME, SPARSE_MODEL_NAME global EMBEDDING_NUM_THREADS, EMBEDDING_BATCH_SIZE, LITELLM_URL, LITELLM_API_KEY + global VISION_OCR_MODEL, CHAT_MODEL cfg = _get_db_emb_config() @@ -52,6 +57,8 @@ def init_embeddings( EMBEDDING_BATCH_SIZE = int(batch_size if batch_size is not None else cfg.get("batch_size", 32)) LITELLM_URL = (litellm_url or cfg.get("litellm_url") or "http://litellm:4000/v1").strip() LITELLM_API_KEY = (litellm_api_key or cfg.get("litellm_api_key") or "dummy").strip() + VISION_OCR_MODEL = (vision_ocr_model or cfg.get("vision_ocr_model") or os.getenv("VISION_OCR_MODEL", "gemini-2.5-flash")).strip() + CHAT_MODEL = (chat_model or cfg.get("chat_model") or os.getenv("CHAT_MODEL", "gemini-2.5-flash")).strip() # Set underlying OpenMP / BLAS thread guard os.environ["OMP_NUM_THREADS"] = str(EMBEDDING_NUM_THREADS) @@ -103,6 +110,8 @@ def get_embedding_config() -> Dict[str, Any]: "threads": EMBEDDING_NUM_THREADS, "batch_size": EMBEDDING_BATCH_SIZE, "litellm_url": LITELLM_URL, + "vision_ocr_model": VISION_OCR_MODEL, + "chat_model": CHAT_MODEL, "system_cpus": sys_res["cpus"], "system_memory_gb": sys_res["memory_gb"], } @@ -115,6 +124,8 @@ def update_embedding_config( batch_size: Optional[int] = None, litellm_url: Optional[str] = None, litellm_api_key: Optional[str] = None, + vision_ocr_model: Optional[str] = None, + chat_model: Optional[str] = None, ) -> Dict[str, Any]: """Updates embedding configuration in SQLite and hot-reloads models in memory.""" from app.services.database import set_embedding_db_config @@ -126,6 +137,8 @@ def update_embedding_config( batch_size=batch_size, litellm_url=litellm_url, litellm_api_key=litellm_api_key, + vision_ocr_model=vision_ocr_model, + chat_model=chat_model, ) init_embeddings( provider=provider, @@ -135,6 +148,8 @@ def update_embedding_config( batch_size=batch_size, litellm_url=litellm_url, litellm_api_key=litellm_api_key, + vision_ocr_model=vision_ocr_model, + chat_model=chat_model, ) return get_embedding_config() diff --git a/app/services/litellm_service.py b/app/services/litellm_service.py new file mode 100644 index 0000000..692056e --- /dev/null +++ b/app/services/litellm_service.py @@ -0,0 +1,182 @@ +import os +import logging +from typing import Optional, Dict, Any, List +import httpx + +from app.services.database import get_embedding_db_config + +logger = logging.getLogger("contextcortex.litellm") + + +async def discover_models( + url: Optional[str] = None, + api_key: Optional[str] = None, +) -> Dict[str, Any]: + """ + Queries LiteLLM GET /v1/models endpoint, categorizes returned models by capability + (embedding, vision OCR, chat completion), and returns structured model metadata. + """ + db_cfg: Dict[str, Any] = {} + try: + db_cfg = get_embedding_db_config() + except Exception as e: + logger.debug(f"Failed to fetch db config for litellm discovery: {e}") + + # Resolve URL: explicitly passed -> SQLite db config -> environment -> fallback + raw_url = ( + (url.strip() if url and url.strip() else None) + or (db_cfg.get("litellm_url") if db_cfg and db_cfg.get("litellm_url") else None) + or os.getenv("LITELLM_URL") + or "http://litellm:4000/v1" + ) + + # Resolve API Key: explicitly passed -> SQLite db config -> environment -> fallback + resolved_api_key = ( + (api_key.strip() if api_key and api_key.strip() else None) + or (db_cfg.get("litellm_api_key") if db_cfg and db_cfg.get("litellm_api_key") else None) + or os.getenv("LITELLM_API_KEY") + or "dummy" + ) + + # Normalize URL to target /models endpoint + clean_url = raw_url.strip().rstrip("/") + if clean_url.endswith("/models"): + endpoint = clean_url + else: + if not clean_url.endswith("/v1"): + clean_url = f"{clean_url}/v1" + endpoint = f"{clean_url}/models" + + headers = {"Authorization": f"Bearer {resolved_api_key}"} + + try: + async with httpx.AsyncClient(timeout=6.0) as client: + response = await client.get(endpoint, headers=headers) + + if response.status_code != 200: + error_msg = f"LiteLLM returned status {response.status_code}: {response.text.strip() or response.reason_phrase}" + logger.warning(f"LiteLLM model discovery failed: {error_msg}") + return { + "status": "error", + "message": error_msg, + "total_models": 0, + "models": [], + "embedding_models": [], + "vision_models": [], + "chat_models": [], + } + + payload = response.json() + raw_items: List[Any] = [] + if isinstance(payload, dict): + if "data" in payload and isinstance(payload["data"], list): + raw_items = payload["data"] + elif "models" in payload and isinstance(payload["models"], list): + raw_items = payload["models"] + elif isinstance(payload, list): + raw_items = payload + + models: List[Dict[str, Any]] = [] + embedding_models: List[str] = [] + vision_models: List[str] = [] + chat_models: List[str] = [] + + vision_patterns = ["vision", "-vl", "flash", "pro", "gemini", "gpt-4", "claude", "qwen3-vl"] + embedding_patterns = ["embed", "bge", "text-embedding"] + image_gen_patterns = ["dall-e", "midjourney", "stable-diffusion", "flux"] + + for item in raw_items: + if isinstance(item, dict): + model_id = str(item.get("id") or item.get("model") or item.get("name") or "").strip() + mode = str(item.get("mode") or "").strip() + normalized_model = dict(item) + elif isinstance(item, str): + model_id = item.strip() + mode = "" + normalized_model = {"id": model_id, "mode": mode} + else: + continue + + if not model_id: + continue + + models.append(normalized_model) + mid_lower = model_id.lower() + mode_lower = mode.lower() + + # Classification + is_embedding = (mode_lower == "embedding") or any(pat in mid_lower for pat in embedding_patterns) + is_image_gen = (mode_lower in ["image_generation", "image-generation", "dall-e"]) or any( + pat in mid_lower for pat in image_gen_patterns + ) + + if is_embedding: + embedding_models.append(model_id) + + if not is_embedding and not is_image_gen: + # Vision / Multimodal model categorization + if mode_lower in ["vision", "multimodal"] or any(pat in mid_lower for pat in vision_patterns): + vision_models.append(model_id) + + # Chat model categorization + if mode_lower == "chat" or mode_lower not in ["embedding", "image_generation"]: + chat_models.append(model_id) + + embedding_models = sorted(list(set(embedding_models))) + vision_models = sorted(list(set(vision_models))) + chat_models = sorted(list(set(chat_models))) + models.sort(key=lambda m: str(m.get("id", ""))) + + return { + "status": "success", + "total_models": len(models), + "models": models, + "embedding_models": embedding_models, + "vision_models": vision_models, + "chat_models": chat_models, + } + + except (httpx.TimeoutException, TimeoutError) as e: + logger.warning(f"LiteLLM model discovery timed out: {e}") + return { + "status": "error", + "message": f"LiteLLM request timed out: {e}", + "total_models": 0, + "models": [], + "embedding_models": [], + "vision_models": [], + "chat_models": [], + } + except httpx.ConnectError as e: + logger.warning(f"LiteLLM connection error: {e}") + return { + "status": "error", + "message": f"LiteLLM connection failed: {e}", + "total_models": 0, + "models": [], + "embedding_models": [], + "vision_models": [], + "chat_models": [], + } + except httpx.HTTPStatusError as e: + logger.warning(f"LiteLLM HTTP error: {e}") + return { + "status": "error", + "message": f"LiteLLM returned HTTP status {e.response.status_code}: {e}", + "total_models": 0, + "models": [], + "embedding_models": [], + "vision_models": [], + "chat_models": [], + } + except Exception as e: + logger.exception(f"Unexpected error discovering LiteLLM models: {e}") + return { + "status": "error", + "message": f"Failed to discover models: {e}", + "total_models": 0, + "models": [], + "embedding_models": [], + "vision_models": [], + "chat_models": [], + } diff --git a/app/services/pdf_extractor.py b/app/services/pdf_extractor.py index 3ccedab..f8660d7 100644 --- a/app/services/pdf_extractor.py +++ b/app/services/pdf_extractor.py @@ -6,6 +6,7 @@ from typing import List, Dict, Any, Union, Optional import pymupdf from openai import OpenAI +from app.services.database import get_vision_ocr_model logger = logging.getLogger("contextcortex.pdf") @@ -34,7 +35,7 @@ def to_dict(self) -> Dict[str, Any]: return asdict(self) -def _call_vision_ocr(png_bytes: bytes, model_name: str = "gemini-2.0-flash") -> str: +def _call_vision_ocr(png_bytes: bytes, model_name: Optional[str] = None) -> str: """Invokes LiteLLM / OpenAI compatible vision model to transcribe document page.""" litellm_url = os.getenv("LITELLM_URL", "http://litellm:4000/v1").strip() litellm_key = os.getenv("LITELLM_API_KEY", "sk-default").strip() @@ -48,8 +49,10 @@ def _call_vision_ocr(png_bytes: bytes, model_name: str = "gemini-2.0-flash") -> "Preserve list structures and code blocks where applicable. Do not summarize or extrapolate." ) + active_model = model_name or get_vision_ocr_model() + response = client.chat.completions.create( - model=os.getenv("VISION_OCR_MODEL", model_name), + model=active_model, messages=[ {"role": "system", "content": system_prompt}, { diff --git a/docs/superpowers/plans/2026-09-07-litellm-model-settings-and-discovery.md b/docs/superpowers/plans/2026-09-07-litellm-model-settings-and-discovery.md new file mode 100644 index 0000000..d4bd1a5 --- /dev/null +++ b/docs/superpowers/plans/2026-09-07-litellm-model-settings-and-discovery.md @@ -0,0 +1,110 @@ +# LiteLLM Model Settings & Dynamic Discovery Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Move all model settings for LiteLLM (embeddings, vision AI OCR, chat completion) into the ContextCortex UI with dynamic model discovery (`GET /admin/api/models/discover`), categorized model selection dropdowns, custom manual overrides, and SQLite persistent storage. + +**Architecture:** +1. Expand SQLite `system_metadata` schema and accessors in `app/services/database/connection.py` to persist `vision_ocr_model` and `chat_model` alongside existing embedding parameters. +2. Build `app/services/litellm_service.py` to query LiteLLM's `GET /v1/models` endpoint, categorize models by capability (Embedding, Vision/OCR, Chat), and handle offline/timeout states gracefully. +3. Expose `GET /admin/api/models/discover` and update `/admin/api/settings/embedding` in `app/api/routers/settings.py`. +4. Modernize `frontend/src/components/settings/EmbeddingSettings.tsx` and `frontend/src/Settings.tsx` with a dynamic "Discover Models" trigger, model dropdowns with instant custom switching, connection badges, and responsive inputs. + +**Tech Stack:** Python 3.12, FastAPI, `httpx`, Pydantic v2, SQLite, React 18, Vite, TypeScript, Vitest, Pytest. + +--- + +## Global Constraints + +- Never break existing FastEmbed local embedding mode defaults (`BAAI/bge-small-en-v1.5`, `Qdrant/bm25`). +- Ensure all model configurations gracefully fall back to environment variables (`LITELLM_URL`, `LITELLM_API_KEY`, `VISION_OCR_MODEL`, `EMBEDDING_MODEL`) if SQLite records are empty. +- Keep network calls to LiteLLM resilient with strict timeouts (<= 6 seconds) to prevent UI or API blocking. +- 100% test pass rate across backend (`pytest`) and frontend (`vitest`), plus successful production build (`npm run build`). + +--- + +## Task Decomposition + +### Task 1: SQLite Metadata Persistence & Dynamic Model Retrieval + +- [ ] Write failing unit test `tests/test_model_metadata_persistence.py` testing `get_embedding_db_config()`, `set_embedding_db_config()`, and `get_vision_ocr_model()` persistence. +- [ ] Run pytest to verify test failure: `pytest tests/test_model_metadata_persistence.py`. +- [ ] In `app/services/database/connection.py`: + - Add `vision_ocr_model` and `chat_model` to `_resolve_default_embedding_config()`, `get_embedding_db_config()`, and `set_embedding_db_config()`. + - Add and export `get_vision_ocr_model() -> str` and `get_chat_model() -> str`. +- [ ] In `app/services/database/__init__.py`: + - Export `get_vision_ocr_model` and `get_chat_model`. +- [ ] In `app/services/embeddings.py`: + - Update `get_embedding_config()` and `update_embedding_config()` to accept and return `vision_ocr_model` and `chat_model`. +- [ ] In `app/services/pdf_extractor.py`: + - Update `_call_vision_ocr` to use `get_vision_ocr_model()` dynamically. +- [ ] In `app/models/schemas.py`: + - Add `vision_ocr_model: Optional[str] = None` and `chat_model: Optional[str] = None` to `EmbeddingSettingsRequest`. +- [ ] Run pytest to verify test passes: `pytest tests/test_model_metadata_persistence.py tests/test_pdf_extractor.py`. +- [ ] Commit: `git commit -m "feat(models): add vision_ocr_model and chat_model persistence to system_metadata"` + +--- + +### Task 2: LiteLLM Model Discovery Service (`app/services/litellm_service.py`) + +- [ ] Write failing unit test `tests/test_litellm_service.py` testing `discover_models` with mock LiteLLM API responses: + - Successful response with mixed models (chat, embedding, vision) + - Classification check: `gemini-embedding-2` in `embedding_models`, `gemini-2.5-flash` in `vision_models` and `chat_models` + - Connection refused / timeout handling returning structured fallback. +- [ ] Run pytest to verify test failure: `pytest tests/test_litellm_service.py`. +- [ ] Implement `app/services/litellm_service.py`: + - Function `async def discover_models(url: Optional[str] = None, api_key: Optional[str] = None) -> Dict[str, Any]` + - Use `httpx.AsyncClient(timeout=6.0)` to request `{url}/models` with `Authorization: Bearer {api_key}` header. + - Categorization algorithm for `embedding_models`, `vision_models`, `chat_models`. + - Graceful exception trapping for `httpx.TimeoutException`, `httpx.ConnectError`, `httpx.HTTPStatusError`. +- [ ] Run pytest to verify test passes: `pytest tests/test_litellm_service.py`. +- [ ] Commit: `git commit -m "feat(litellm): add dynamic model discovery service with classification and resilient fallbacks"` + +--- + +### Task 3: Backend API Routes for Model Discovery and Unified Settings + +- [ ] Write failing test `tests/test_model_discovery_api.py` testing: + - `GET /admin/api/models/discover` + - `GET /admin/api/settings/embedding` (returns `vision_ocr_model` and `chat_model`) + - `POST /admin/api/settings/embedding` (saves and updates `vision_ocr_model` and `chat_model`) +- [ ] Run pytest to verify failure: `pytest tests/test_model_discovery_api.py`. +- [ ] In `app/api/routers/settings.py`: + - Add route `@router.get("/admin/api/models/discover")` calling `litellm_service.discover_models`. + - Update `api_save_embedding_settings` to pass `payload.vision_ocr_model` and `payload.chat_model` to `emb_service.update_embedding_config`. +- [ ] Run pytest to verify test passes: `pytest tests/test_model_discovery_api.py`. +- [ ] Commit: `git commit -m "feat(api): expose /admin/api/models/discover and update embedding settings route"` + +--- + +### Task 4: Frontend UI for Model Discovery & Dynamic Selection + +- [ ] In `frontend/src/types.ts`: + - Extend `EmbeddingConfig` with `vision_ocr_model?: string` and `chat_model?: string`. + - Add `ModelDiscoveryResult` interface. +- [ ] In `frontend/src/components/settings/EmbeddingSettings.tsx`: + - Add "Discover Models" button next to LiteLLM URL and API Key inputs. + - When models are discovered, display count badge (e.g. `42 models detected`). + - Render dynamic ` + {/* LiteLLM Specific Connection & Discovery Controls */} + {embProvider === 'api' && ( + <> +
+
+ + setEmbLitellmApiKey?.(e.target.value)} + placeholder="sk-..." + autoComplete="off" + /> +
+ +
+ +
+
+ + {/* Discovery Feedback Banner / Status */} + {discoveryResult && ( +
+ {discoveryResult.status === 'success' ? ( +
+ + + Connected to LiteLLM — {discoveryResult.total_models} models available ({embeddingModels.length} embedding, {visionModels.length} vision, {chatModels.length} chat) + +
+ ) : ( +
+ + {discoveryResult.message || 'Could not connect to LiteLLM endpoint'} +
+ )} +
+ )} + + {/* Dense Model Selection with Discovery Dropdown */} +
+
+ + {embeddingModels.length > 0 && ( + + )} +
+ + {!customDense && embeddingModels.length > 0 ? ( + + ) : ( + setEmbDenseModel(e.target.value)} + placeholder="gemini-embedding-2" + /> + )} +
+ + )} + + {/* Local Sparse BM25 Model */} {embProvider === 'local' && (
@@ -189,6 +324,98 @@ export function EmbeddingSettings({
)} + {/* Vision AI OCR Model Selection */} +
+
+ + {visionModels.length > 0 && ( + + )} +
+ + {!customVision && visionModels.length > 0 ? ( + + ) : ( + setEmbVisionOcrModel?.(e.target.value)} + placeholder="gemini-2.5-flash" + /> + )} +
+ + {/* General Chat / Completion Model Selection */} +
+
+ + {chatModels.length > 0 && ( + + )} +
+ + {!customChat && chatModels.length > 0 ? ( + + ) : ( + setEmbChatModel?.(e.target.value)} + placeholder="gemini-2.5-flash" + /> + )} +
+
diff --git a/frontend/src/tests/EmbeddingSettings.test.tsx b/frontend/src/tests/EmbeddingSettings.test.tsx index 7f94e95..ca39f3d 100644 --- a/frontend/src/tests/EmbeddingSettings.test.tsx +++ b/frontend/src/tests/EmbeddingSettings.test.tsx @@ -141,4 +141,213 @@ describe('EmbeddingSettings Component', () => { fireEvent.click(saveBtn); expect(onSave).toHaveBeenCalled(); }); + + it('renders model discovery controls and triggers onDiscoverModels when button clicked', () => { + const onDiscover = vi.fn(); + const setApiKey = vi.fn(); + + render( + + ); + + const apiKeyInput = screen.getByLabelText(/LiteLLM API Key/i); + expect(apiKeyInput).toBeInTheDocument(); + fireEvent.change(apiKeyInput, { target: { value: 'sk-new-key' } }); + expect(setApiKey).toHaveBeenCalledWith('sk-new-key'); + + const discoverBtn = screen.getByRole('button', { name: /Discover Models/i }); + fireEvent.click(discoverBtn); + expect(onDiscover).toHaveBeenCalled(); + }); + + it('renders discovered model dropdowns and allows selecting models', () => { + const setDense = vi.fn(); + const setVision = vi.fn(); + const setChat = vi.fn(); + + const mockDiscovery = { + status: 'success' as const, + total_models: 4, + models: [ + { id: 'gemini-embedding-2', mode: 'embedding' }, + { id: 'text-embedding-3-small', mode: 'embedding' }, + { id: 'gemini-2.5-flash', mode: 'chat' }, + { id: 'qwen3-vl-32b-instruct', mode: 'chat' } + ], + embedding_models: ['gemini-embedding-2', 'text-embedding-3-small'], + vision_models: ['gemini-2.5-flash', 'qwen3-vl-32b-instruct'], + chat_models: ['gemini-2.5-flash', 'qwen3-vl-32b-instruct'] + }; + + render( + + ); + + expect(screen.getByText(/4 models available/i)).toBeInTheDocument(); + + const denseSelect = screen.getByLabelText(/^Dense Embedding Model/i); + fireEvent.change(denseSelect, { target: { value: 'text-embedding-3-small' } }); + expect(setDense).toHaveBeenCalledWith('text-embedding-3-small'); + + const visionSelect = screen.getByLabelText(/Vision AI OCR Model/i); + fireEvent.change(visionSelect, { target: { value: 'qwen3-vl-32b-instruct' } }); + expect(setVision).toHaveBeenCalledWith('qwen3-vl-32b-instruct'); + + const chatSelect = screen.getByLabelText(/General Chat & Synthesis Model/i); + fireEvent.change(chatSelect, { target: { value: 'qwen3-vl-32b-instruct' } }); + expect(setChat).toHaveBeenCalledWith('qwen3-vl-32b-instruct'); + }); + + it('displays discovery error banner when LiteLLM is unreachable', () => { + render( + + ); + + expect(screen.getByText(/Connection to http:\/\/invalid:4000\/v1 timed out/i)).toBeInTheDocument(); + }); + + it('switches to manual input when Custom is selected from dropdown or link clicked', () => { + const setDense = vi.fn(); + const mockDiscovery = { + status: 'success' as const, + total_models: 2, + models: [ + { id: 'gemini-embedding-2', mode: 'embedding' }, + { id: 'text-embedding-3-small', mode: 'embedding' } + ], + embedding_models: ['gemini-embedding-2', 'text-embedding-3-small'], + vision_models: [], + chat_models: [] + }; + + render( + + ); + + // Click "Enter custom model →" button + const customToggleBtn = screen.getByRole('button', { name: /Enter custom model/i }); + expect(customToggleBtn).toBeInTheDocument(); + fireEvent.click(customToggleBtn); + + // Should now be a text input instead of select + const manualInput = screen.getByPlaceholderText('gemini-embedding-2'); + expect(manualInput.tagName.toLowerCase()).toBe('input'); + fireEvent.change(manualInput, { target: { value: 'my-custom-model-id' } }); + expect(setDense).toHaveBeenCalledWith('my-custom-model-id'); + }); + + it('disables discover button and shows spinner while isDiscovering is true', () => { + render( + + ); + + const discoverBtn = screen.getByRole('button', { name: /Discovering\.\.\./i }); + expect(discoverBtn).toBeDisabled(); + expect(screen.getByText(/Discovering\.\.\./i)).toBeInTheDocument(); + }); }); + diff --git a/frontend/src/tests/Settings.test.tsx b/frontend/src/tests/Settings.test.tsx index 12559cb..e9713e2 100644 --- a/frontend/src/tests/Settings.test.tsx +++ b/frontend/src/tests/Settings.test.tsx @@ -44,7 +44,9 @@ const mockEmbeddingConfig: EmbeddingConfig = { batch_size: 32, system_cpus: 8, system_memory_gb: 16.0, - litellm_url: 'http://litellm:4000/v1' + litellm_url: 'http://litellm:4000/v1', + vision_ocr_model: 'gemini-2.5-flash', + chat_model: 'gemini-2.5-flash', }; const mockHostCreds: GitHostCredential[] = [ @@ -1118,7 +1120,9 @@ describe('Settings Component', () => { threads: 4, batch_size: 64, dense_model: 'BAAI/bge-small-en-v1.5', - sparse_model: 'Qdrant/bm25' + sparse_model: 'Qdrant/bm25', + vision_ocr_model: 'gemini-2.5-flash', + chat_model: 'gemini-2.5-flash' }) }) ); diff --git a/frontend/src/types.ts b/frontend/src/types.ts index 7da1ce8..395cbb0 100644 --- a/frontend/src/types.ts +++ b/frontend/src/types.ts @@ -39,10 +39,32 @@ export interface EmbeddingConfig { batch_size: number; litellm_url?: string; litellm_api_key?: string; + vision_ocr_model?: string; + chat_model?: string; system_cpus?: number; system_memory_gb?: number; } +export interface DiscoveredModel { + id: string; + object?: string; + created?: number; + owned_by?: string; + mode?: string; + max_input_tokens?: number; + max_output_tokens?: number; +} + +export interface ModelDiscoveryResult { + status: 'success' | 'error'; + total_models: number; + models: DiscoveredModel[]; + embedding_models: string[]; + vision_models: string[]; + chat_models: string[]; + message?: string; +} + export interface AutoSyncSettings { interval_mins: number; webhook_url: string; diff --git a/tests/test_litellm_service.py b/tests/test_litellm_service.py new file mode 100644 index 0000000..5d510b3 --- /dev/null +++ b/tests/test_litellm_service.py @@ -0,0 +1,185 @@ +import os +import pytest +import httpx +from unittest.mock import AsyncMock, patch, MagicMock + +from app.services.litellm_service import discover_models + + +@pytest.mark.asyncio +async def test_discover_models_success(): + mock_payload = { + "data": [ + {"id": "gemini-embedding-2", "mode": "embedding", "owned_by": "openai"}, + {"id": "text-embedding-3-small", "mode": "embedding", "owned_by": "openai"}, + {"id": "bge-m3", "mode": "embedding", "owned_by": "local"}, + {"id": "gemini-2.5-flash", "mode": "chat", "owned_by": "google"}, + {"id": "qwen3-vl-32b-instruct", "mode": "chat", "owned_by": "alibaba"}, + {"id": "gemini-2.5-pro", "mode": "chat", "owned_by": "google"}, + {"id": "deepseek-v3.2", "mode": "chat", "owned_by": "deepseek"}, + {"id": "dall-e-3", "mode": "image_generation", "owned_by": "openai"}, + ], + "object": "list", + } + + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.json.return_value = mock_payload + mock_response.raise_for_status = MagicMock() + + with patch("httpx.AsyncClient.get", new_callable=AsyncMock) as mock_get: + mock_get.return_value = mock_response + + res = await discover_models( + url="http://litellm-test:4000/v1", + api_key="sk-test-key", + ) + + assert res["status"] == "success" + assert res["total_models"] == 8 + assert len(res["models"]) == 8 + + # Embedding models check + assert "gemini-embedding-2" in res["embedding_models"] + assert "text-embedding-3-small" in res["embedding_models"] + assert "bge-m3" in res["embedding_models"] + assert "gemini-2.5-flash" not in res["embedding_models"] + + # Vision models check + assert "gemini-2.5-flash" in res["vision_models"] + assert "qwen3-vl-32b-instruct" in res["vision_models"] + assert "gemini-2.5-pro" in res["vision_models"] + assert "deepseek-v3.2" not in res["vision_models"] + assert "gemini-embedding-2" not in res["vision_models"] + assert "dall-e-3" not in res["vision_models"] + + # Chat models check + assert "gemini-2.5-flash" in res["chat_models"] + assert "gemini-2.5-pro" in res["chat_models"] + assert "deepseek-v3.2" in res["chat_models"] + assert "qwen3-vl-32b-instruct" in res["chat_models"] + assert "gemini-embedding-2" not in res["chat_models"] + assert "dall-e-3" not in res["chat_models"] + + # Sorting check + assert res["embedding_models"] == sorted(res["embedding_models"]) + assert res["vision_models"] == sorted(res["vision_models"]) + assert res["chat_models"] == sorted(res["chat_models"]) + + # Request check + mock_get.assert_called_once() + args, kwargs = mock_get.call_args + assert args[0] == "http://litellm-test:4000/v1/models" + assert kwargs["headers"]["Authorization"] == "Bearer sk-test-key" + + +@pytest.mark.asyncio +async def test_discover_models_timeout(): + with patch("httpx.AsyncClient.get", new_callable=AsyncMock) as mock_get: + mock_get.side_effect = httpx.TimeoutException("Connection timed out after 6.0s") + + res = await discover_models(url="http://litellm-timeout:4000/v1", api_key="sk-test") + + assert res["status"] == "error" + assert "timed out" in res["message"].lower() or "timeout" in res["message"].lower() + assert res["models"] == [] + assert res["embedding_models"] == [] + assert res["vision_models"] == [] + assert res["chat_models"] == [] + + +@pytest.mark.asyncio +async def test_discover_models_connect_error(): + with patch("httpx.AsyncClient.get", new_callable=AsyncMock) as mock_get: + mock_get.side_effect = httpx.ConnectError("Failed to establish connection") + + res = await discover_models(url="http://litellm-unreachable:4000/v1", api_key="sk-test") + + assert res["status"] == "error" + assert "connect" in res["message"].lower() or "connection" in res["message"].lower() + assert res["models"] == [] + assert res["embedding_models"] == [] + assert res["vision_models"] == [] + assert res["chat_models"] == [] + + +@pytest.mark.asyncio +async def test_discover_models_http_401_unauthorized(): + mock_request = httpx.Request("GET", "http://litellm:4000/v1/models") + mock_response = httpx.Response(status_code=401, request=mock_request, text="Unauthorized: Invalid API Key") + + with patch("httpx.AsyncClient.get", new_callable=AsyncMock) as mock_get: + mock_get.return_value = mock_response + + res = await discover_models(url="http://litellm:4000/v1", api_key="bad-key") + + assert res["status"] == "error" + assert "401" in res["message"] or "unauthorized" in res["message"].lower() + assert res["models"] == [] + assert res["embedding_models"] == [] + assert res["vision_models"] == [] + assert res["chat_models"] == [] + + +@pytest.mark.asyncio +async def test_discover_models_http_500_error(): + mock_request = httpx.Request("GET", "http://litellm:4000/v1/models") + mock_response = httpx.Response(status_code=500, request=mock_request, text="Internal Server Error") + + with patch("httpx.AsyncClient.get", new_callable=AsyncMock) as mock_get: + mock_get.return_value = mock_response + + res = await discover_models(url="http://litellm:4000/v1", api_key="test-key") + + assert res["status"] == "error" + assert "500" in res["message"] + assert res["models"] == [] + + +@pytest.mark.asyncio +async def test_discover_models_url_normalization(): + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.json.return_value = {"data": []} + mock_response.raise_for_status = MagicMock() + + with patch("httpx.AsyncClient.get", new_callable=AsyncMock) as mock_get: + mock_get.return_value = mock_response + + # Test with trailing slashes and without /v1 + await discover_models(url="http://litellm:4000/", api_key="dummy") + args, _ = mock_get.call_args + assert args[0] == "http://litellm:4000/v1/models" + + # Test with /v1/ trailing slash + await discover_models(url="http://litellm:4000/v1/", api_key="dummy") + args, _ = mock_get.call_args + assert args[0] == "http://litellm:4000/v1/models" + + # Test with already full /models path + await discover_models(url="http://litellm:4000/v1/models", api_key="dummy") + args, _ = mock_get.call_args + assert args[0] == "http://litellm:4000/v1/models" + + +@pytest.mark.asyncio +async def test_discover_models_default_resolution(monkeypatch): + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.json.return_value = {"data": [{"id": "model-1", "mode": "chat"}]} + mock_response.raise_for_status = MagicMock() + + with patch("httpx.AsyncClient.get", new_callable=AsyncMock) as mock_get: + mock_get.return_value = mock_response + + with patch("app.services.litellm_service.get_embedding_db_config") as mock_db_cfg: + mock_db_cfg.return_value = { + "litellm_url": "http://custom-db-litellm:4000/v1", + "litellm_api_key": "db-secret-key", + } + + res = await discover_models() + assert res["status"] == "success" + args, kwargs = mock_get.call_args + assert args[0] == "http://custom-db-litellm:4000/v1/models" + assert kwargs["headers"]["Authorization"] == "Bearer db-secret-key" diff --git a/tests/test_model_discovery_api.py b/tests/test_model_discovery_api.py new file mode 100644 index 0000000..e6977be --- /dev/null +++ b/tests/test_model_discovery_api.py @@ -0,0 +1,79 @@ +import pytest +from unittest.mock import patch, AsyncMock +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from app.api.routes import router + +app = FastAPI() +app.include_router(router) +client = TestClient(app) + + +def test_api_discover_models_endpoint(): + mock_discovery_result = { + "status": "success", + "total_models": 3, + "models": [ + {"id": "gemini-embedding-2", "mode": "embedding"}, + {"id": "gemini-2.5-flash", "mode": "chat"}, + {"id": "qwen3-vl-32b-instruct", "mode": "chat"}, + ], + "embedding_models": ["gemini-embedding-2"], + "vision_models": ["gemini-2.5-flash", "qwen3-vl-32b-instruct"], + "chat_models": ["gemini-2.5-flash", "qwen3-vl-32b-instruct"], + } + + with patch("app.services.litellm_service.discover_models", new_callable=AsyncMock, return_value=mock_discovery_result): + resp = client.get("/admin/api/models/discover?url=http://custom:4000/v1&api_key=sk-custom") + assert resp.status_code == 200 + data = resp.json() + assert data["status"] == "success" + assert "gemini-embedding-2" in data["embedding_models"] + assert "gemini-2.5-flash" in data["vision_models"] + + +def test_api_embedding_settings_get_and_post_with_models(tmp_path, monkeypatch): + test_db = str(tmp_path / "test_models.db") + monkeypatch.setattr("app.services.database.connection.CACHE_DB_PATH", test_db) + monkeypatch.setattr("app.services.database.CACHE_DB_PATH", test_db) + + from app.services.database.connection import init_db + init_db() + + # 1. GET embedding settings + get_resp = client.get("/admin/api/settings/embedding") + assert get_resp.status_code == 200 + cfg = get_resp.json() + assert "vision_ocr_model" in cfg + assert "chat_model" in cfg + + # 2. POST embedding settings with vision_ocr_model and chat_model + post_payload = { + "provider": "api", + "dense_model": "gemini-embedding-2", + "sparse_model": "Qdrant/bm25", + "threads": 4, + "batch_size": 64, + "litellm_url": "http://litellm:4000/v1", + "litellm_api_key": "sk-test-key", + "vision_ocr_model": "gemini-2.5-flash", + "chat_model": "gemini-2.5-pro", + } + + post_resp = client.post("/admin/api/settings/embedding", json=post_payload) + assert post_resp.status_code == 200 + res = post_resp.json() + assert res["status"] == "success" + saved = res["config"] + assert saved["provider"] == "api" + assert saved["dense_model"] == "gemini-embedding-2" + assert saved["vision_ocr_model"] == "gemini-2.5-flash" + assert saved["chat_model"] == "gemini-2.5-pro" + + # 3. Verify GET returns updated configuration + verify_resp = client.get("/admin/api/settings/embedding") + assert verify_resp.status_code == 200 + vcfg = verify_resp.json() + assert vcfg["vision_ocr_model"] == "gemini-2.5-flash" + assert vcfg["chat_model"] == "gemini-2.5-pro" diff --git a/tests/test_model_metadata_persistence.py b/tests/test_model_metadata_persistence.py new file mode 100644 index 0000000..065f729 --- /dev/null +++ b/tests/test_model_metadata_persistence.py @@ -0,0 +1,93 @@ +import os +import pytest +from app.services.database.engine import get_db_engine, init_db +from app.services.database.connection import ( + get_metadata, + get_embedding_db_config, + set_embedding_db_config, +) + +def setup_test_db(tmp_path, monkeypatch): + db_file = tmp_path / "test_models.db" + db_url = f"sqlite:///{db_file}" + monkeypatch.setenv("DATABASE_URL", db_url) + engine = get_db_engine(db_url, reset=True) + init_db(engine=engine) + return engine + + +def test_embedding_db_config_defaults(tmp_path, monkeypatch): + # Ensure env vars are cleared + monkeypatch.delenv("VISION_OCR_MODEL", raising=False) + monkeypatch.delenv("CHAT_MODEL", raising=False) + setup_test_db(tmp_path, monkeypatch) + + from app.services.database import get_vision_ocr_model, get_chat_model + + cfg = get_embedding_db_config() + assert "vision_ocr_model" in cfg + assert "chat_model" in cfg + assert cfg["vision_ocr_model"] == "gemini-2.5-flash" + assert cfg["chat_model"] == "gemini-2.5-flash" + + assert get_vision_ocr_model() == "gemini-2.5-flash" + assert get_chat_model() == "gemini-2.5-flash" + + +def test_env_variable_fallbacks(tmp_path, monkeypatch): + monkeypatch.setenv("VISION_OCR_MODEL", "custom-ocr-env") + monkeypatch.setenv("CHAT_MODEL", "custom-chat-env") + setup_test_db(tmp_path, monkeypatch) + + from app.services.database import get_vision_ocr_model, get_chat_model + + cfg = get_embedding_db_config() + assert cfg["vision_ocr_model"] == "custom-ocr-env" + assert cfg["chat_model"] == "custom-chat-env" + + assert get_vision_ocr_model() == "custom-ocr-env" + assert get_chat_model() == "custom-chat-env" + + +def test_set_embedding_db_config_persists_to_system_metadata(tmp_path, monkeypatch): + monkeypatch.setenv("VISION_OCR_MODEL", "fallback-vision") + monkeypatch.setenv("CHAT_MODEL", "fallback-chat") + setup_test_db(tmp_path, monkeypatch) + + from app.services.database import get_vision_ocr_model, get_chat_model + + set_embedding_db_config( + vision_ocr_model="db-persisted-vision", + chat_model="db-persisted-chat", + ) + + # Check system_metadata persistence + assert get_metadata("vision_ocr_model") == "db-persisted-vision" or get_metadata("embedding_vision_ocr_model") == "db-persisted-vision" + assert get_metadata("chat_model") == "db-persisted-chat" or get_metadata("embedding_chat_model") == "db-persisted-chat" + + # Check config retrieval + cfg = get_embedding_db_config() + assert cfg["vision_ocr_model"] == "db-persisted-vision" + assert cfg["chat_model"] == "db-persisted-chat" + + # Check accessor functions prioritize DB over environment + assert get_vision_ocr_model() == "db-persisted-vision" + assert get_chat_model() == "db-persisted-chat" + + +def test_update_embedding_config_service(tmp_path, monkeypatch): + setup_test_db(tmp_path, monkeypatch) + + from app.services.embeddings import update_embedding_config, get_embedding_config + + updated = update_embedding_config( + vision_ocr_model="service-ocr-model", + chat_model="service-chat-model", + ) + + assert updated.get("vision_ocr_model") == "service-ocr-model" + assert updated.get("chat_model") == "service-chat-model" + + active_cfg = get_embedding_config() + assert active_cfg.get("vision_ocr_model") == "service-ocr-model" + assert active_cfg.get("chat_model") == "service-chat-model"