Commit ed273297 authored by Kantz's avatar Kantz
Browse files

llm provider umgestellt

parent bd1e3cb2
......@@ -2,33 +2,42 @@ MATHPIX_APP_ID=""
MATHPIX_APP_KEY=""
POSTGRES_URL=""
DAILY_LLM_CALL_LIMIT="100"
DAILY_LLM_TOKEN_LIMIT="50000"
DAILY_LLM_CALL_LIMIT="400"
DAILY_LLM_TOKEN_LIMIT="500000"
FRONTEND_URL="http://localhost:5173"
FRONTEND_URL="http://frontend:3000"
ORCHESTRATOR="tutor" # "tutor" or "qa"
ORCHESTRATOR="task" # "tutor", "task" or "qa"
RETRIEVAL_IMPL="child" # "child" or "subsection"
LLM_PROVIDER="openai" # "openai", "mistral", or "ollama"
LLM_PROVIDER="gwdg" # "openai", "gwdg", "mistral", or "ollama"
OPENAI_BASE_URL="https://chat-ai.academiccloud.de/v1/"
EMBEDDING_PROVIDER="gwdg" # "sentence-transformer", "openai", or "gwdg"
EMBEDDING_TYPE="sentence-transformer" # Legacy/internal fallback: "openai-like" or "sentence-transformer"
EMBEDDING_DIM="512"
SENTENCE_TRANSFORMER_MODEL="jinaai/jina-embeddings-v5-text-small-retrieval"
GWDG_BASE_URL="https://chat-ai.academiccloud.de/v1/"
GWDG_API_KEY=""
GWDG_CHAT_MODEL="glm-4.7"
GWDG_CHAT_TEMPERATURE="0.2"
GWDG_EMBED_MODEL="e5-mistral-7b-instruct"
GWDG_TIMEOUT="60"
OPENAI_BASE_URL="https://api.openai.com/v1/"
OPENAI_API_KEY=""
OPENAI_CHAT_MODEL="mistral-large-3-675b-instruct-2512"
OPENAI_CHAT_MODEL="gpt-5.4-nano"
OPENAI_CHAT_TEMPERATURE="0.2"
OPENAI_EMBED_MODEL="e5-mistral-7b-instruct"
OPENAI_EMBED_MODEL="text-embedding-3-small"
OPENAI_TIMEOUT="60"
EMBEDDING_TYPE="sentence-transformer" # "openai-like" or "sentence-transformer"
EMBEDDING_DIM="512"
SENTENCE_TRANSFORMER_MODEL="jinaai/jina-embeddings-v4"
OLLAMA_URL="http://localhost:11434"
OLLAMA_MODEL= "ministral-3"
OLLAMA_URL=""
OLLAMA_MODEL= "gemma4:26b"
OLLAMA_TEMPERATURE="0.2"
OLLAMA_TIMEOUT="60"
OLLAMA_TIMEOUT="120"
MISTRAL_CHAT_MODEL="mistral-large-3-675b-instruct-2512"
MISTRAL_API_KEY=""
MISTRAL_CHAT_TIMEOUT="60"
MISTRAL_CHAT_TEMPERATURE="0.2"
MISTRAL_TIMEOUT="60"
......@@ -7,6 +7,7 @@ from typing import Any, Dict
import httpx
import psycopg
from fastapi import APIRouter
from fastapi.responses import JSONResponse
import app.config as config
from app.deterministic_services import llm_quota
......@@ -66,8 +67,7 @@ def _normalize_openai_models_url(base_url: str) -> str:
return f"{trimmed}/v1/models"
def _check_openai() -> dict:
settings = config.get_openai_base_settings()
def _check_openai_compatible(settings: config.OpenAIBaseSettings | None) -> dict:
if not settings:
return {"status": "missing_config"}
......@@ -84,6 +84,39 @@ def _check_openai() -> dict:
return {"status": "error", "url": url, "detail": str(exc)}
def _check_openai() -> dict:
return _check_openai_compatible(config.get_openai_base_settings())
def _check_gwdg() -> dict:
return _check_openai_compatible(config.get_gwdg_base_settings())
def _check_mistral() -> dict:
try:
settings = config.get_mistral_chat_settings()
except ValueError as exc:
return {"status": "missing_config", "detail": str(exc)}
if not settings:
return {"status": "missing_config"}
return {"status": "ok", "model": settings.model}
def _check_selected_llm_provider() -> dict[str, dict]:
try:
provider = config.get_llm_provider()
except ValueError as exc:
return {"llm_provider": {"status": "missing_config", "detail": str(exc)}}
checks = {
"openai": _check_openai,
"gwdg": _check_gwdg,
"mistral": _check_mistral,
"ollama": _check_ollama,
}
return {provider: checks[provider]()}
def _check_postgres() -> dict:
try:
pg_url = config.get_postgres_url()
......@@ -125,8 +158,7 @@ def _check_llm_quota() -> dict:
@router.get("/api/health")
def health() -> Dict[str, Any]:
services = {
"ollama": _check_ollama(),
"openai": _check_openai(),
**_check_selected_llm_provider(),
"postgres": _check_postgres(),
"llm_quota": _check_llm_quota(),
}
......@@ -137,12 +169,12 @@ def health() -> Dict[str, Any]:
return {"status": overall, "services": services}
@router.get("/api/health/ready")
def readiness() -> Dict[str, Any]:
@router.get("/api/health/ready", response_model=None)
def readiness() -> Any:
state = get_readiness_state()
if state.get("status") == "ready":
return state
return {"status_code": 503, "content": state}
return JSONResponse(status_code=503, content=state)
def run_startup_checks() -> Dict[str, Any]:
......
......@@ -7,7 +7,8 @@ from typing import Optional
load_dotenv()
SUPPORTED_LLM_PROVIDERS = {"openai", "mistral", "ollama"}
SUPPORTED_LLM_PROVIDERS = {"openai", "gwdg", "mistral", "ollama"}
SUPPORTED_EMBEDDING_PROVIDERS = {"sentence-transformer", "openai", "gwdg"}
class EmbeddingSettings(BaseModel):
......@@ -44,31 +45,68 @@ def get_llm_provider() -> str:
return provider
def get_embedding_provider() -> str:
value = os.getenv("EMBEDDING_PROVIDER")
if value is not None and value.strip():
provider = value.strip().lower()
if provider in SUPPORTED_EMBEDDING_PROVIDERS:
return provider
supported = ", ".join(sorted(SUPPORTED_EMBEDDING_PROVIDERS))
raise ValueError(
f"Unsupported EMBEDDING_PROVIDER: {value}. Expected one of: {supported}"
)
legacy_type = os.getenv("EMBEDDING_TYPE", "openai-like").strip().lower()
if legacy_type == "sentence-transformer":
return "sentence-transformer"
if legacy_type == "openai-like":
return "openai"
supported = ", ".join(sorted(SUPPORTED_EMBEDDING_PROVIDERS))
raise ValueError(
f"Unsupported EMBEDDING_TYPE: {legacy_type}. Set EMBEDDING_PROVIDER to one of: {supported}"
)
def get_embedding_settings() -> EmbeddingSettings:
embedding_type = os.getenv("EMBEDDING_TYPE", "openai-like")
if embedding_type == "sentence-transformer":
provider = get_embedding_provider()
if provider == "sentence-transformer":
return EmbeddingSettings(
embedding_type=embedding_type,
embedding_type="sentence-transformer",
model=os.getenv("SENTENCE_TRANSFORMER_MODEL",
"jinaai/jina-embeddings-v5-text-small-retrieval"),
target_dim=int(os.getenv("EMBEDDING_DIM", "1024")),
)
if embedding_type == "openai-like":
model_target_dim = int(os.getenv("EMBEDDING_DIM", "1024"))
if provider == "openai":
base_url = os.getenv("OPENAI_BASE_URL")
api_key = os.getenv("OPENAI_API_KEY")
model = os.getenv("OPENAI_EMBED_MODEL", "e5-mistral-7b-instruct")
model_target_dim = int(os.getenv("EMBEDDING_DIM", "1024"))
model = os.getenv("OPENAI_EMBED_MODEL", "text-embedding-3-small")
if not base_url or not api_key:
raise ValueError("Missing OPENAI_BASE_URL or OPENAI_API_KEY")
raise ValueError("Missing OPENAI_BASE_URL or OPENAI_API_KEY for embeddings")
return EmbeddingSettings(
embedding_type=embedding_type,
embedding_type="openai-like",
base_url=base_url,
api_key=api_key,
model=model,
target_dim=model_target_dim,
)
else:
raise ValueError(f"Unsupported EMBEDDING_TYPE: {embedding_type}")
if provider == "gwdg":
base_url = os.getenv("GWDG_BASE_URL")
api_key = os.getenv("GWDG_API_KEY")
model = os.getenv("GWDG_EMBED_MODEL")
if not base_url or not api_key or not model:
raise ValueError("Missing GWDG_BASE_URL, GWDG_API_KEY or GWDG_EMBED_MODEL")
return EmbeddingSettings(
embedding_type="openai-like",
base_url=base_url,
api_key=api_key,
model=model,
target_dim=model_target_dim,
)
raise ValueError(f"Unsupported EMBEDDING_PROVIDER: {provider}")
@dataclass(frozen=True)
......@@ -168,6 +206,17 @@ def get_openai_base_settings() -> OpenAIBaseSettings | None:
)
def get_gwdg_base_settings() -> OpenAIBaseSettings | None:
base_url = os.getenv("GWDG_BASE_URL")
api_key = os.getenv("GWDG_API_KEY")
if not base_url or not api_key:
return None
return OpenAIBaseSettings(
base_url=base_url,
api_key=api_key,
)
def get_openai_chat_settings() -> OpenAIChatSettings | None:
model = os.getenv("OPENAI_CHAT_MODEL")
if not model:
......@@ -179,11 +228,27 @@ def get_openai_chat_settings() -> OpenAIChatSettings | None:
base_url=base_settings.base_url,
api_key=base_settings.api_key,
model=model,
timeout=_read_float(os.getenv("OPENAI_CHAT_TIMEOUT")),
timeout=_read_float(os.getenv("OPENAI_TIMEOUT") or os.getenv("OPENAI_CHAT_TIMEOUT")),
temperature=_read_float(os.getenv("OPENAI_CHAT_TEMPERATURE")),
)
def get_gwdg_chat_settings() -> OpenAIChatSettings | None:
model = os.getenv("GWDG_CHAT_MODEL")
if not model:
return None
base_settings = get_gwdg_base_settings()
if not base_settings:
raise ValueError("Missing GWDG_BASE_URL or GWDG_API_KEY for chat")
return OpenAIChatSettings(
base_url=base_settings.base_url,
api_key=base_settings.api_key,
model=model,
timeout=_read_float(os.getenv("GWDG_TIMEOUT")),
temperature=_read_float(os.getenv("GWDG_CHAT_TEMPERATURE")),
)
def get_mistral_chat_settings() -> MistralChatSettings | None:
model = os.getenv("MISTRAL_CHAT_MODEL")
api_key = os.getenv("MISTRAL_API_KEY")
......@@ -194,7 +259,7 @@ def get_mistral_chat_settings() -> MistralChatSettings | None:
return MistralChatSettings(
api_key=api_key,
model=model,
timeout=_read_float(os.getenv("MISTRAL_CHAT_TIMEOUT")),
timeout=_read_float(os.getenv("MISTRAL_TIMEOUT") or os.getenv("MISTRAL_CHAT_TIMEOUT")),
temperature=_read_float(os.getenv("MISTRAL_CHAT_TEMPERATURE")),
)
......
......@@ -4,7 +4,7 @@ import math
from typing import List, Optional, Union
from enum import Enum
import httpx
from openai import OpenAI
from pydantic import BaseModel, Field
from sentence_transformers import SentenceTransformer
......@@ -27,6 +27,7 @@ class OpenAILikeConfig(BaseModel):
model: str = Field(...,
description="Modellname (z. B. 'nomic-embed-text')")
target_dim: int = Field(1024, description="Ziel-Dimension der Embeddings")
timeout: float | None = Field(None, description="Request timeout in seconds")
class SentenceTransformerConfig(BaseModel):
......@@ -99,43 +100,31 @@ class OpenAILikeEmbeddings(BaseEmbeddings):
self.base_url = config.base_url.rstrip("/")
self.api_key = config.api_key
self.model = config.model
self.endpoint = self._embedding_endpoint()
def _embedding_endpoint(self) -> str:
"""Berechnet den korrekten Endpunkt für die Embedding-API."""
if self.base_url.endswith("/embeddings"):
return self.base_url
if self.base_url.endswith("/v1"):
return f"{self.base_url}/embeddings"
return f"{self.base_url}/v1/embeddings"
self.timeout = config.timeout or 60.0
def _embed(self, inputs: List[str] | str) -> List[List[float]]:
"""Ruft die externe Embedding-API auf."""
payload = {
"input": inputs,
"model": self.model,
"encoding_format": "float",
}
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}",
}
with httpx.Client(timeout=60.0) as client:
response = client.post(
self.endpoint, headers=headers, json=payload)
response.raise_for_status()
data = response.json().get("data")
client = OpenAI(
api_key=self.api_key,
base_url=self.base_url,
timeout=self.timeout,
)
response = client.embeddings.create(
input=inputs,
model=self.model,
encoding_format="float",
)
data = response.get("data") if isinstance(response, dict) else getattr(response, "data", None)
if not isinstance(data, list):
raise ValueError("Embedding response missing 'data' list.")
# Sortiere nach Index, falls nötig
data_sorted = sorted(data, key=lambda item: item.get("index", 0))
data_sorted = sorted(data, key=_embedding_item_index)
embeddings: List[List[float]] = []
for item in data_sorted:
emb = item.get("embedding")
emb = _embedding_item_vector(item)
if not isinstance(emb, list):
raise ValueError("Embedding item missing 'embedding' list.")
embeddings.append(self._truncate([float(x) for x in emb]))
......@@ -143,6 +132,18 @@ class OpenAILikeEmbeddings(BaseEmbeddings):
return embeddings
def _embedding_item_index(item: object) -> int:
if isinstance(item, dict):
return int(item.get("index", 0) or 0)
return int(getattr(item, "index", 0) or 0)
def _embedding_item_vector(item: object) -> object:
if isinstance(item, dict):
return item.get("embedding")
return getattr(item, "embedding", None)
class SentenceTransformerEmbeddings(BaseEmbeddings):
"""
Embeddings-Wrapper für lokale SentenceTransformer Modelle.
......
......@@ -113,6 +113,16 @@ def _require_openai_chat_settings() -> config.OpenAIChatSettings:
return settings
def _require_gwdg_chat_settings() -> config.OpenAIChatSettings:
settings = config.get_gwdg_chat_settings()
if not settings:
raise ValueError(
"LLM_PROVIDER=gwdg requires GWDG_CHAT_MODEL, GWDG_BASE_URL, "
"and GWDG_API_KEY"
)
return settings
def _require_mistral_chat_settings() -> config.MistralChatSettings:
settings = config.get_mistral_chat_settings()
if not settings:
......@@ -122,7 +132,7 @@ def _require_mistral_chat_settings() -> config.MistralChatSettings:
return settings
def _chat_openai(
def _chat_openai_compatible(
messages: list[dict],
settings: config.OpenAIChatSettings | None = None,
) -> dict:
......@@ -207,10 +217,18 @@ def chat(
_warn_deprecated_provider_flags(use_ollama, use_mistral)
provider = config.get_llm_provider()
if tools and provider != "ollama":
raise RuntimeError(
"Deprecated LLM toolcalling is only implemented for LLM_PROVIDER=ollama. "
f"Current LLM_PROVIDER={provider}."
)
if provider == "openai":
settings = _require_openai_chat_settings()
return _quota_tracked_chat(lambda: _chat_openai(messages, settings))
return _quota_tracked_chat(lambda: _chat_openai_compatible(messages, settings))
if provider == "gwdg":
settings = _require_gwdg_chat_settings()
return _quota_tracked_chat(lambda: _chat_openai_compatible(messages, settings))
if provider == "mistral":
settings = _require_mistral_chat_settings()
return _quota_tracked_chat(lambda: _chat_mistral(messages, settings))
......
import os
import unittest
from unittest.mock import patch
os.environ.setdefault("EMBEDDING_PROVIDER", "sentence-transformer")
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
from app.deterministic_services.orchestrators.orchestrator_base import (
ChatState,
finalize_response,
......
import argparse
import json
import os
from typing import Any
os.environ.setdefault("EMBEDDING_PROVIDER", "sentence-transformer")
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
from app.LLM_services import decision_LLM
from app.deterministic_services import context_store
......
......@@ -3,6 +3,7 @@ from __future__ import annotations
import importlib.util
import os
import sys
from types import SimpleNamespace
from pathlib import Path
import unittest
from unittest.mock import MagicMock, patch
......@@ -26,12 +27,15 @@ config = _load_module("backend_config_test_module", "app/config.py")
fake_sentence_transformers = type(sys)("sentence_transformers")
fake_sentence_transformers.SentenceTransformer = object
sys.modules.setdefault("sentence_transformers", fake_sentence_transformers)
fake_openai = type(sys)("openai")
fake_openai.OpenAI = object
sys.modules.setdefault("openai", fake_openai)
embeddings = _load_module("backend_embeddings_test_module", "app/deterministic_services/embeddings.py")
class SentenceTransformerJinaV5Test(unittest.TestCase):
def test_config_defaults_to_jina_v5(self) -> None:
with patch.dict(os.environ, {"EMBEDDING_TYPE": "sentence-transformer"}, clear=False):
with patch.dict(os.environ, {"EMBEDDING_TYPE": "sentence-transformer"}, clear=True):
settings = config.get_embedding_settings()
self.assertEqual(settings.model, "jinaai/jina-embeddings-v5-text-small-retrieval")
......@@ -61,6 +65,57 @@ class SentenceTransformerJinaV5Test(unittest.TestCase):
self.assertEqual(len(docs[0]), 4)
self.assertEqual(len(query), 4)
def test_openai_like_embedder_uses_openai_library(self) -> None:
class FakeEmbeddingsClient:
create_kwargs: dict | None = None
def create(self, **kwargs):
FakeEmbeddingsClient.create_kwargs = kwargs
return SimpleNamespace(
data=[
SimpleNamespace(index=1, embedding=[0.0, 3.0, 4.0]),
SimpleNamespace(index=0, embedding=[3.0, 4.0, 0.0]),
]
)
class FakeOpenAI:
init_kwargs: dict | None = None
def __init__(self, **kwargs):
FakeOpenAI.init_kwargs = kwargs
self.embeddings = FakeEmbeddingsClient()
with patch.object(embeddings, "OpenAI", FakeOpenAI):
embedder = embeddings.OpenAILikeEmbeddings(
embeddings.OpenAILikeConfig(
base_url="https://chat-ai.academiccloud.de/v1/",
api_key="gwdg-key",
model="e5-mistral-7b-instruct",
target_dim=2,
timeout=12.5,
)
)
result = embedder.embed_documents(["a", "b"])
self.assertEqual(
FakeOpenAI.init_kwargs,
{
"api_key": "gwdg-key",
"base_url": "https://chat-ai.academiccloud.de/v1",
"timeout": 12.5,
},
)
self.assertEqual(
FakeEmbeddingsClient.create_kwargs,
{
"input": ["a", "b"],
"model": "e5-mistral-7b-instruct",
"encoding_format": "float",
},
)
self.assertEqual(len(result), 2)
self.assertEqual(len(result[0]), 2)
if __name__ == "__main__":
unittest.main()
import json
import os
import unittest
from unittest.mock import patch
os.environ.setdefault("EMBEDDING_PROVIDER", "sentence-transformer")
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
from app.api import health
from fastapi.responses import JSONResponse
......@@ -43,6 +48,35 @@ class HealthReadinessUnitTest(unittest.TestCase):
},
)
def test_health_checks_only_selected_gwdg_provider(self) -> None:
with patch.dict(os.environ, {"LLM_PROVIDER": "gwdg"}), patch(
"app.api.health._check_gwdg",
return_value={"status": "ok", "url": "https://chat-ai.academiccloud.de/v1/models"},
) as gwdg_check, patch(
"app.api.health._check_openai"
) as openai_check, patch(
"app.api.health._check_ollama"
) as ollama_check, patch(
"app.api.health._check_mistral"
) as mistral_check, patch(
"app.api.health._check_postgres",
return_value={"status": "ok"},
), patch(
"app.api.health._check_llm_quota",
return_value={"status": "ok"},
):
response = health.health()
self.assertEqual(response["status"], "ok")
self.assertEqual(
set(response["services"].keys()),
{"gwdg", "postgres", "llm_quota"},
)
gwdg_check.assert_called_once()
openai_check.assert_not_called()
ollama_check.assert_not_called()
mistral_check.assert_not_called()
if __name__ == "__main__":
unittest.main()
import argparse
import json
import os
from typing import Any
from app.LLM_services import hint_LLM
os.environ.setdefault("EMBEDDING_PROVIDER", "sentence-transformer")
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
from app.LLM_services import open_hint_LLM as hint_LLM
from app.deterministic_services import context_store
......
import argparse
import json
import os
from typing import Iterable
os.environ.setdefault("EMBEDDING_PROVIDER", "sentence-transformer")
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
from app.LLM_services import math_intent_LLM
......
import argparse
import os
os.environ.setdefault("EMBEDDING_PROVIDER", "sentence-transformer")
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
from app import config
from app.deterministic_services.embeddings import EmbeddingFactory
......
......@@ -81,7 +81,7 @@ def _dummy_tool() -> str:
class LLMProviderConfigTest(unittest.TestCase):
def test_get_llm_provider_accepts_supported_values(self) -> None:
for provider in ("openai", "mistral", "ollama"):
for provider in ("openai", "gwdg", "mistral", "ollama"):
with self.subTest(provider=provider), patch.dict(
os.environ, {"LLM_PROVIDER": provider}, clear=True
):
......@@ -101,6 +101,58 @@ class LLMProviderConfigTest(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "Unsupported LLM_PROVIDER"):
config.get_llm_provider()
def test_get_gwdg_chat_settings_reads_gwdg_keys(self) -> None:
env = {
"GWDG_BASE_URL": "https://chat-ai.academiccloud.de/v1/",
"GWDG_API_KEY": "gwdg-key",
"GWDG_CHAT_MODEL": "glm-4.7",
"GWDG_CHAT_TEMPERATURE": "0.2",
"GWDG_TIMEOUT": "60",
}
with patch.dict(os.environ, env, clear=True):
settings = config.get_gwdg_chat_settings()
self.assertIsNotNone(settings)
assert settings is not None
self.assertEqual(settings.base_url, env["GWDG_BASE_URL"])
self.assertEqual(settings.api_key, env["GWDG_API_KEY"])
self.assertEqual(settings.model, env["GWDG_CHAT_MODEL"])
self.assertEqual(settings.temperature, 0.2)
self.assertEqual(settings.timeout, 60.0)
def test_get_embedding_provider_accepts_supported_values(self) -> None:
for provider in ("sentence-transformer", "openai", "gwdg"):
with self.subTest(provider=provider), patch.dict(
os.environ, {"EMBEDDING_PROVIDER": provider}, clear=True
):
self.assertEqual(config.get_embedding_provider(), provider)
def test_get_embedding_provider_uses_legacy_embedding_type_fallback(self) -> None:
with patch.dict(os.environ, {"EMBEDDING_TYPE": "openai-like"}, clear=True):
self.assertEqual(config.get_embedding_provider(), "openai")
with patch.dict(os.environ, {"EMBEDDING_TYPE": "sentence-transformer"}, clear=True):
self.assertEqual(config.get_embedding_provider(), "sentence-transformer")
def test_get_embedding_settings_reads_gwdg_keys(self) -> None:
env = {
"EMBEDDING_PROVIDER": "gwdg",
"EMBEDDING_DIM": "512",
"GWDG_BASE_URL": "https://chat-ai.academiccloud.de/v1/",
"GWDG_API_KEY": "gwdg-key",
"GWDG_EMBED_MODEL": "e5-mistral-7b-instruct",
"GWDG_TIMEOUT": "60",
}
with patch.dict(os.environ, env, clear=True):
settings = config.get_embedding_settings()
self.assertEqual(settings.embedding_type, "openai-like")
self.assertEqual(settings.base_url, env["GWDG_BASE_URL"])
self.assertEqual(settings.api_key, env["GWDG_API_KEY"])
self.assertEqual(settings.model, env["GWDG_EMBED_MODEL"])
self.assertEqual(settings.target_dim, 512)
self.assertEqual(settings.timeout, 60.0)
class LLMClientProviderTest(unittest.TestCase):
def test_chat_uses_only_openai_provider(self) -> None:
......@@ -113,7 +165,7 @@ class LLMClientProviderTest(unittest.TestCase):
), patch.object(
llm_client, "_record_call", side_effect=lambda result, tokens=None: result
), patch.object(
llm_client, "_chat_openai", return_value=expected
llm_client, "_chat_openai_compatible", return_value=expected
) as openai_chat, patch.object(
llm_client, "_chat_mistral"
) as mistral_chat, patch.object(
......@@ -126,6 +178,29 @@ class LLMClientProviderTest(unittest.TestCase):
mistral_chat.assert_not_called()
ollama_chat.assert_not_called()
def test_chat_uses_only_gwdg_provider(self) -> None:
settings = object()
expected = {"raw": object(), "message": {"content": "gwdg"}}
with patch.dict(os.environ, {"LLM_PROVIDER": "gwdg"}), patch.object(
llm_client, "_require_gwdg_chat_settings", return_value=settings
), patch.object(
llm_client, "_ensure_within_llm_quota"
), patch.object(
llm_client, "_record_call", side_effect=lambda result, tokens=None: result
), patch.object(
llm_client, "_chat_openai_compatible", return_value=expected
) as compatible_chat, patch.object(
llm_client, "_chat_mistral"
) as mistral_chat, patch.object(
llm_client, "_chat_ollama"
) as ollama_chat:
result = llm_client.chat(MESSAGES)
self.assertEqual(result, expected)
compatible_chat.assert_called_once_with(MESSAGES, settings)
mistral_chat.assert_not_called()
ollama_chat.assert_not_called()
def test_chat_uses_only_mistral_provider(self) -> None:
settings = object()
expected = {"raw": object(), "message": {"content": "mistral"}}
......@@ -136,7 +211,7 @@ class LLMClientProviderTest(unittest.TestCase):
), patch.object(
llm_client, "_record_call", side_effect=lambda result, tokens=None: result
), patch.object(
llm_client, "_chat_openai"
llm_client, "_chat_openai_compatible"
) as openai_chat, patch.object(
llm_client, "_chat_mistral", return_value=expected
) as mistral_chat, patch.object(
......@@ -152,7 +227,7 @@ class LLMClientProviderTest(unittest.TestCase):
def test_chat_uses_only_ollama_provider(self) -> None:
expected = {"raw": object(), "message": {"content": "ollama"}}
with patch.dict(os.environ, {"LLM_PROVIDER": "ollama"}), patch.object(
llm_client, "_chat_openai"
llm_client, "_chat_openai_compatible"
) as openai_chat, patch.object(
llm_client, "_chat_mistral"
) as mistral_chat, patch.object(
......@@ -175,7 +250,7 @@ class LLMClientProviderTest(unittest.TestCase):
), patch.object(
llm_client, "_record_call", side_effect=lambda result, tokens=None: result
), patch.object(
llm_client, "_chat_openai", return_value=expected
llm_client, "_chat_openai_compatible", return_value=expected
) as openai_chat, patch.object(
llm_client, "_chat_ollama"
) as ollama_chat:
......@@ -186,6 +261,15 @@ class LLMClientProviderTest(unittest.TestCase):
openai_chat.assert_called_once_with(MESSAGES, settings)
ollama_chat.assert_not_called()
def test_selected_gwdg_config_error_happens_before_quota(self) -> None:
with patch.dict(os.environ, {"LLM_PROVIDER": "gwdg"}, clear=True), patch.object(
llm_client, "_ensure_within_llm_quota"
) as ensure_quota:
with self.assertRaisesRegex(ValueError, "LLM_PROVIDER=gwdg requires"):
llm_client.chat(MESSAGES)
ensure_quota.assert_not_called()
def test_selected_provider_config_error_happens_before_quota(self) -> None:
with patch.dict(os.environ, {"LLM_PROVIDER": "openai"}, clear=True), patch.object(
llm_client, "_ensure_within_llm_quota"
......@@ -197,7 +281,7 @@ class LLMClientProviderTest(unittest.TestCase):
def test_chat_without_provider_does_not_fallback(self) -> None:
with patch.dict(os.environ, {}, clear=True), patch.object(
llm_client, "_chat_openai"
llm_client, "_chat_openai_compatible"
) as openai_chat, patch.object(
llm_client, "_chat_mistral"
) as mistral_chat, patch.object(
......
from __future__ import annotations
import os
import unittest
os.environ.setdefault("EMBEDDING_PROVIDER", "sentence-transformer")
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
from app.deterministic_services.vector_store import (
Retrieved,
Source,
......
Supports Markdown
0% or .
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment