Commit ed273297 authored by Kantz's avatar Kantz
Browse files

llm provider umgestellt

parent bd1e3cb2
...@@ -2,33 +2,42 @@ MATHPIX_APP_ID="" ...@@ -2,33 +2,42 @@ MATHPIX_APP_ID=""
MATHPIX_APP_KEY="" MATHPIX_APP_KEY=""
POSTGRES_URL="" POSTGRES_URL=""
DAILY_LLM_CALL_LIMIT="100" DAILY_LLM_CALL_LIMIT="400"
DAILY_LLM_TOKEN_LIMIT="50000" 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" 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_API_KEY=""
OPENAI_CHAT_MODEL="mistral-large-3-675b-instruct-2512" OPENAI_CHAT_MODEL="gpt-5.4-nano"
OPENAI_CHAT_TEMPERATURE="0.2" OPENAI_CHAT_TEMPERATURE="0.2"
OPENAI_EMBED_MODEL="e5-mistral-7b-instruct" OPENAI_EMBED_MODEL="text-embedding-3-small"
OPENAI_TIMEOUT="60" OPENAI_TIMEOUT="60"
EMBEDDING_TYPE="sentence-transformer" # "openai-like" or "sentence-transformer" OLLAMA_URL=""
EMBEDDING_DIM="512" OLLAMA_MODEL= "gemma4:26b"
SENTENCE_TRANSFORMER_MODEL="jinaai/jina-embeddings-v4"
OLLAMA_URL="http://localhost:11434"
OLLAMA_MODEL= "ministral-3"
OLLAMA_TEMPERATURE="0.2" OLLAMA_TEMPERATURE="0.2"
OLLAMA_TIMEOUT="60" OLLAMA_TIMEOUT="120"
MISTRAL_CHAT_MODEL="mistral-large-3-675b-instruct-2512" MISTRAL_CHAT_MODEL="mistral-large-3-675b-instruct-2512"
MISTRAL_API_KEY="" MISTRAL_API_KEY=""
MISTRAL_CHAT_TIMEOUT="60" MISTRAL_CHAT_TIMEOUT="60"
MISTRAL_CHAT_TEMPERATURE="0.2" MISTRAL_CHAT_TEMPERATURE="0.2"
MISTRAL_TIMEOUT="60"
...@@ -7,6 +7,7 @@ from typing import Any, Dict ...@@ -7,6 +7,7 @@ from typing import Any, Dict
import httpx import httpx
import psycopg import psycopg
from fastapi import APIRouter from fastapi import APIRouter
from fastapi.responses import JSONResponse
import app.config as config import app.config as config
from app.deterministic_services import llm_quota from app.deterministic_services import llm_quota
...@@ -66,8 +67,7 @@ def _normalize_openai_models_url(base_url: str) -> str: ...@@ -66,8 +67,7 @@ def _normalize_openai_models_url(base_url: str) -> str:
return f"{trimmed}/v1/models" return f"{trimmed}/v1/models"
def _check_openai() -> dict: def _check_openai_compatible(settings: config.OpenAIBaseSettings | None) -> dict:
settings = config.get_openai_base_settings()
if not settings: if not settings:
return {"status": "missing_config"} return {"status": "missing_config"}
...@@ -84,6 +84,39 @@ def _check_openai() -> dict: ...@@ -84,6 +84,39 @@ def _check_openai() -> dict:
return {"status": "error", "url": url, "detail": str(exc)} 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: def _check_postgres() -> dict:
try: try:
pg_url = config.get_postgres_url() pg_url = config.get_postgres_url()
...@@ -125,8 +158,7 @@ def _check_llm_quota() -> dict: ...@@ -125,8 +158,7 @@ def _check_llm_quota() -> dict:
@router.get("/api/health") @router.get("/api/health")
def health() -> Dict[str, Any]: def health() -> Dict[str, Any]:
services = { services = {
"ollama": _check_ollama(), **_check_selected_llm_provider(),
"openai": _check_openai(),
"postgres": _check_postgres(), "postgres": _check_postgres(),
"llm_quota": _check_llm_quota(), "llm_quota": _check_llm_quota(),
} }
...@@ -137,12 +169,12 @@ def health() -> Dict[str, Any]: ...@@ -137,12 +169,12 @@ def health() -> Dict[str, Any]:
return {"status": overall, "services": services} return {"status": overall, "services": services}
@router.get("/api/health/ready") @router.get("/api/health/ready", response_model=None)
def readiness() -> Dict[str, Any]: def readiness() -> Any:
state = get_readiness_state() state = get_readiness_state()
if state.get("status") == "ready": if state.get("status") == "ready":
return state return state
return {"status_code": 503, "content": state} return JSONResponse(status_code=503, content=state)
def run_startup_checks() -> Dict[str, Any]: def run_startup_checks() -> Dict[str, Any]:
......
...@@ -7,7 +7,8 @@ from typing import Optional ...@@ -7,7 +7,8 @@ from typing import Optional
load_dotenv() 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): class EmbeddingSettings(BaseModel):
...@@ -44,31 +45,68 @@ def get_llm_provider() -> str: ...@@ -44,31 +45,68 @@ def get_llm_provider() -> str:
return provider 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: def get_embedding_settings() -> EmbeddingSettings:
embedding_type = os.getenv("EMBEDDING_TYPE", "openai-like") provider = get_embedding_provider()
if embedding_type == "sentence-transformer": if provider == "sentence-transformer":
return EmbeddingSettings( return EmbeddingSettings(
embedding_type=embedding_type, embedding_type="sentence-transformer",
model=os.getenv("SENTENCE_TRANSFORMER_MODEL", model=os.getenv("SENTENCE_TRANSFORMER_MODEL",
"jinaai/jina-embeddings-v5-text-small-retrieval"), "jinaai/jina-embeddings-v5-text-small-retrieval"),
target_dim=int(os.getenv("EMBEDDING_DIM", "1024")), 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") base_url = os.getenv("OPENAI_BASE_URL")
api_key = os.getenv("OPENAI_API_KEY") api_key = os.getenv("OPENAI_API_KEY")
model = os.getenv("OPENAI_EMBED_MODEL", "e5-mistral-7b-instruct") model = os.getenv("OPENAI_EMBED_MODEL", "text-embedding-3-small")
model_target_dim = int(os.getenv("EMBEDDING_DIM", "1024"))
if not base_url or not api_key: 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( return EmbeddingSettings(
embedding_type=embedding_type, embedding_type="openai-like",
base_url=base_url, base_url=base_url,
api_key=api_key, api_key=api_key,
model=model, model=model,
target_dim=model_target_dim, 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) @dataclass(frozen=True)
...@@ -168,6 +206,17 @@ def get_openai_base_settings() -> OpenAIBaseSettings | None: ...@@ -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: def get_openai_chat_settings() -> OpenAIChatSettings | None:
model = os.getenv("OPENAI_CHAT_MODEL") model = os.getenv("OPENAI_CHAT_MODEL")
if not model: if not model:
...@@ -179,11 +228,27 @@ def get_openai_chat_settings() -> OpenAIChatSettings | None: ...@@ -179,11 +228,27 @@ def get_openai_chat_settings() -> OpenAIChatSettings | None:
base_url=base_settings.base_url, base_url=base_settings.base_url,
api_key=base_settings.api_key, api_key=base_settings.api_key,
model=model, 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")), 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: def get_mistral_chat_settings() -> MistralChatSettings | None:
model = os.getenv("MISTRAL_CHAT_MODEL") model = os.getenv("MISTRAL_CHAT_MODEL")
api_key = os.getenv("MISTRAL_API_KEY") api_key = os.getenv("MISTRAL_API_KEY")
...@@ -194,7 +259,7 @@ def get_mistral_chat_settings() -> MistralChatSettings | None: ...@@ -194,7 +259,7 @@ def get_mistral_chat_settings() -> MistralChatSettings | None:
return MistralChatSettings( return MistralChatSettings(
api_key=api_key, api_key=api_key,
model=model, 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")), temperature=_read_float(os.getenv("MISTRAL_CHAT_TEMPERATURE")),
) )
......
...@@ -4,7 +4,7 @@ import math ...@@ -4,7 +4,7 @@ import math
from typing import List, Optional, Union from typing import List, Optional, Union
from enum import Enum from enum import Enum
import httpx from openai import OpenAI
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from sentence_transformers import SentenceTransformer from sentence_transformers import SentenceTransformer
...@@ -27,6 +27,7 @@ class OpenAILikeConfig(BaseModel): ...@@ -27,6 +27,7 @@ class OpenAILikeConfig(BaseModel):
model: str = Field(..., model: str = Field(...,
description="Modellname (z. B. 'nomic-embed-text')") description="Modellname (z. B. 'nomic-embed-text')")
target_dim: int = Field(1024, description="Ziel-Dimension der Embeddings") target_dim: int = Field(1024, description="Ziel-Dimension der Embeddings")
timeout: float | None = Field(None, description="Request timeout in seconds")
class SentenceTransformerConfig(BaseModel): class SentenceTransformerConfig(BaseModel):
...@@ -99,43 +100,31 @@ class OpenAILikeEmbeddings(BaseEmbeddings): ...@@ -99,43 +100,31 @@ class OpenAILikeEmbeddings(BaseEmbeddings):
self.base_url = config.base_url.rstrip("/") self.base_url = config.base_url.rstrip("/")
self.api_key = config.api_key self.api_key = config.api_key
self.model = config.model self.model = config.model
self.endpoint = self._embedding_endpoint() self.timeout = config.timeout or 60.0
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"
def _embed(self, inputs: List[str] | str) -> List[List[float]]: def _embed(self, inputs: List[str] | str) -> List[List[float]]:
"""Ruft die externe Embedding-API auf.""" """Ruft die externe Embedding-API auf."""
payload = { client = OpenAI(
"input": inputs, api_key=self.api_key,
"model": self.model, base_url=self.base_url,
"encoding_format": "float", timeout=self.timeout,
} )
headers = { response = client.embeddings.create(
"Content-Type": "application/json", input=inputs,
"Authorization": f"Bearer {self.api_key}", model=self.model,
} encoding_format="float",
)
with httpx.Client(timeout=60.0) as client:
response = client.post( data = response.get("data") if isinstance(response, dict) else getattr(response, "data", None)
self.endpoint, headers=headers, json=payload)
response.raise_for_status()
data = response.json().get("data")
if not isinstance(data, list): if not isinstance(data, list):
raise ValueError("Embedding response missing 'data' list.") raise ValueError("Embedding response missing 'data' list.")
# Sortiere nach Index, falls nötig # 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]] = [] embeddings: List[List[float]] = []
for item in data_sorted: for item in data_sorted:
emb = item.get("embedding") emb = _embedding_item_vector(item)
if not isinstance(emb, list): if not isinstance(emb, list):
raise ValueError("Embedding item missing 'embedding' list.") raise ValueError("Embedding item missing 'embedding' list.")
embeddings.append(self._truncate([float(x) for x in emb])) embeddings.append(self._truncate([float(x) for x in emb]))
...@@ -143,6 +132,18 @@ class OpenAILikeEmbeddings(BaseEmbeddings): ...@@ -143,6 +132,18 @@ class OpenAILikeEmbeddings(BaseEmbeddings):
return embeddings 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): class SentenceTransformerEmbeddings(BaseEmbeddings):
""" """
Embeddings-Wrapper für lokale SentenceTransformer Modelle. Embeddings-Wrapper für lokale SentenceTransformer Modelle.
......
...@@ -113,6 +113,16 @@ def _require_openai_chat_settings() -> config.OpenAIChatSettings: ...@@ -113,6 +113,16 @@ def _require_openai_chat_settings() -> config.OpenAIChatSettings:
return settings 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: def _require_mistral_chat_settings() -> config.MistralChatSettings:
settings = config.get_mistral_chat_settings() settings = config.get_mistral_chat_settings()
if not settings: if not settings:
...@@ -122,7 +132,7 @@ def _require_mistral_chat_settings() -> config.MistralChatSettings: ...@@ -122,7 +132,7 @@ def _require_mistral_chat_settings() -> config.MistralChatSettings:
return settings return settings
def _chat_openai( def _chat_openai_compatible(
messages: list[dict], messages: list[dict],
settings: config.OpenAIChatSettings | None = None, settings: config.OpenAIChatSettings | None = None,
) -> dict: ) -> dict:
...@@ -207,10 +217,18 @@ def chat( ...@@ -207,10 +217,18 @@ def chat(
_warn_deprecated_provider_flags(use_ollama, use_mistral) _warn_deprecated_provider_flags(use_ollama, use_mistral)
provider = config.get_llm_provider() 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": if provider == "openai":
settings = _require_openai_chat_settings() 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": if provider == "mistral":
settings = _require_mistral_chat_settings() settings = _require_mistral_chat_settings()
return _quota_tracked_chat(lambda: _chat_mistral(messages, settings)) return _quota_tracked_chat(lambda: _chat_mistral(messages, settings))
......
import os
import unittest import unittest
from unittest.mock import patch 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 ( from app.deterministic_services.orchestrators.orchestrator_base import (
ChatState, ChatState,
finalize_response, finalize_response,
......
import argparse import argparse
import json import json
import os
from typing import Any 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.LLM_services import decision_LLM
from app.deterministic_services import context_store from app.deterministic_services import context_store
......
...@@ -3,6 +3,7 @@ from __future__ import annotations ...@@ -3,6 +3,7 @@ from __future__ import annotations
import importlib.util import importlib.util
import os import os
import sys import sys
from types import SimpleNamespace
from pathlib import Path from pathlib import Path
import unittest import unittest
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
...@@ -26,12 +27,15 @@ config = _load_module("backend_config_test_module", "app/config.py") ...@@ -26,12 +27,15 @@ config = _load_module("backend_config_test_module", "app/config.py")
fake_sentence_transformers = type(sys)("sentence_transformers") fake_sentence_transformers = type(sys)("sentence_transformers")
fake_sentence_transformers.SentenceTransformer = object fake_sentence_transformers.SentenceTransformer = object
sys.modules.setdefault("sentence_transformers", fake_sentence_transformers) 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") embeddings = _load_module("backend_embeddings_test_module", "app/deterministic_services/embeddings.py")
class SentenceTransformerJinaV5Test(unittest.TestCase): class SentenceTransformerJinaV5Test(unittest.TestCase):
def test_config_defaults_to_jina_v5(self) -> None: 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() settings = config.get_embedding_settings()
self.assertEqual(settings.model, "jinaai/jina-embeddings-v5-text-small-retrieval") self.assertEqual(settings.model, "jinaai/jina-embeddings-v5-text-small-retrieval")
...@@ -61,6 +65,57 @@ class SentenceTransformerJinaV5Test(unittest.TestCase): ...@@ -61,6 +65,57 @@ class SentenceTransformerJinaV5Test(unittest.TestCase):
self.assertEqual(len(docs[0]), 4) self.assertEqual(len(docs[0]), 4)
self.assertEqual(len(query), 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__": if __name__ == "__main__":
unittest.main() unittest.main()
import json import json
import os
import unittest 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 app.api import health
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
...@@ -43,6 +48,35 @@ class HealthReadinessUnitTest(unittest.TestCase): ...@@ -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__": if __name__ == "__main__":
unittest.main() unittest.main()
import argparse import argparse
import json import json
import os
from typing import Any 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 from app.deterministic_services import context_store
......
import argparse import argparse
import json import json
import os
from typing import Iterable 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 from app.LLM_services import math_intent_LLM
......
import argparse import argparse
import os
os.environ.setdefault("EMBEDDING_PROVIDER", "sentence-transformer")
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
from app import config from app import config
from app.deterministic_services.embeddings import EmbeddingFactory from app.deterministic_services.embeddings import EmbeddingFactory
......
...@@ -81,7 +81,7 @@ def _dummy_tool() -> str: ...@@ -81,7 +81,7 @@ def _dummy_tool() -> str:
class LLMProviderConfigTest(unittest.TestCase): class LLMProviderConfigTest(unittest.TestCase):
def test_get_llm_provider_accepts_supported_values(self) -> None: 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( with self.subTest(provider=provider), patch.dict(
os.environ, {"LLM_PROVIDER": provider}, clear=True os.environ, {"LLM_PROVIDER": provider}, clear=True
): ):
...@@ -101,6 +101,58 @@ class LLMProviderConfigTest(unittest.TestCase): ...@@ -101,6 +101,58 @@ class LLMProviderConfigTest(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "Unsupported LLM_PROVIDER"): with self.assertRaisesRegex(ValueError, "Unsupported LLM_PROVIDER"):
config.get_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): class LLMClientProviderTest(unittest.TestCase):
def test_chat_uses_only_openai_provider(self) -> None: def test_chat_uses_only_openai_provider(self) -> None:
...@@ -113,7 +165,7 @@ class LLMClientProviderTest(unittest.TestCase): ...@@ -113,7 +165,7 @@ class LLMClientProviderTest(unittest.TestCase):
), patch.object( ), patch.object(
llm_client, "_record_call", side_effect=lambda result, tokens=None: result llm_client, "_record_call", side_effect=lambda result, tokens=None: result
), patch.object( ), patch.object(
llm_client, "_chat_openai", return_value=expected llm_client, "_chat_openai_compatible", return_value=expected
) as openai_chat, patch.object( ) as openai_chat, patch.object(
llm_client, "_chat_mistral" llm_client, "_chat_mistral"
) as mistral_chat, patch.object( ) as mistral_chat, patch.object(
...@@ -126,6 +178,29 @@ class LLMClientProviderTest(unittest.TestCase): ...@@ -126,6 +178,29 @@ class LLMClientProviderTest(unittest.TestCase):
mistral_chat.assert_not_called() mistral_chat.assert_not_called()
ollama_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: def test_chat_uses_only_mistral_provider(self) -> None:
settings = object() settings = object()
expected = {"raw": object(), "message": {"content": "mistral"}} expected = {"raw": object(), "message": {"content": "mistral"}}
...@@ -136,7 +211,7 @@ class LLMClientProviderTest(unittest.TestCase): ...@@ -136,7 +211,7 @@ class LLMClientProviderTest(unittest.TestCase):
), patch.object( ), patch.object(
llm_client, "_record_call", side_effect=lambda result, tokens=None: result llm_client, "_record_call", side_effect=lambda result, tokens=None: result
), patch.object( ), patch.object(
llm_client, "_chat_openai" llm_client, "_chat_openai_compatible"
) as openai_chat, patch.object( ) as openai_chat, patch.object(
llm_client, "_chat_mistral", return_value=expected llm_client, "_chat_mistral", return_value=expected
) as mistral_chat, patch.object( ) as mistral_chat, patch.object(
...@@ -152,7 +227,7 @@ class LLMClientProviderTest(unittest.TestCase): ...@@ -152,7 +227,7 @@ class LLMClientProviderTest(unittest.TestCase):
def test_chat_uses_only_ollama_provider(self) -> None: def test_chat_uses_only_ollama_provider(self) -> None:
expected = {"raw": object(), "message": {"content": "ollama"}} expected = {"raw": object(), "message": {"content": "ollama"}}
with patch.dict(os.environ, {"LLM_PROVIDER": "ollama"}), patch.object( with patch.dict(os.environ, {"LLM_PROVIDER": "ollama"}), patch.object(
llm_client, "_chat_openai" llm_client, "_chat_openai_compatible"
) as openai_chat, patch.object( ) as openai_chat, patch.object(
llm_client, "_chat_mistral" llm_client, "_chat_mistral"
) as mistral_chat, patch.object( ) as mistral_chat, patch.object(
...@@ -175,7 +250,7 @@ class LLMClientProviderTest(unittest.TestCase): ...@@ -175,7 +250,7 @@ class LLMClientProviderTest(unittest.TestCase):
), patch.object( ), patch.object(
llm_client, "_record_call", side_effect=lambda result, tokens=None: result llm_client, "_record_call", side_effect=lambda result, tokens=None: result
), patch.object( ), patch.object(
llm_client, "_chat_openai", return_value=expected llm_client, "_chat_openai_compatible", return_value=expected
) as openai_chat, patch.object( ) as openai_chat, patch.object(
llm_client, "_chat_ollama" llm_client, "_chat_ollama"
) as ollama_chat: ) as ollama_chat:
...@@ -186,6 +261,15 @@ class LLMClientProviderTest(unittest.TestCase): ...@@ -186,6 +261,15 @@ class LLMClientProviderTest(unittest.TestCase):
openai_chat.assert_called_once_with(MESSAGES, settings) openai_chat.assert_called_once_with(MESSAGES, settings)
ollama_chat.assert_not_called() 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: def test_selected_provider_config_error_happens_before_quota(self) -> None:
with patch.dict(os.environ, {"LLM_PROVIDER": "openai"}, clear=True), patch.object( with patch.dict(os.environ, {"LLM_PROVIDER": "openai"}, clear=True), patch.object(
llm_client, "_ensure_within_llm_quota" llm_client, "_ensure_within_llm_quota"
...@@ -197,7 +281,7 @@ class LLMClientProviderTest(unittest.TestCase): ...@@ -197,7 +281,7 @@ class LLMClientProviderTest(unittest.TestCase):
def test_chat_without_provider_does_not_fallback(self) -> None: def test_chat_without_provider_does_not_fallback(self) -> None:
with patch.dict(os.environ, {}, clear=True), patch.object( with patch.dict(os.environ, {}, clear=True), patch.object(
llm_client, "_chat_openai" llm_client, "_chat_openai_compatible"
) as openai_chat, patch.object( ) as openai_chat, patch.object(
llm_client, "_chat_mistral" llm_client, "_chat_mistral"
) as mistral_chat, patch.object( ) as mistral_chat, patch.object(
......
from __future__ import annotations from __future__ import annotations
import os
import unittest import unittest
os.environ.setdefault("EMBEDDING_PROVIDER", "sentence-transformer")
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
from app.deterministic_services.vector_store import ( from app.deterministic_services.vector_store import (
Retrieved, Retrieved,
Source, 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