Commit 2a0d1d1d authored by Kantz's avatar Kantz
Browse files

Merge branch 'dev' into 'main'

Dev

See merge request kantz/tutor_react!20
parents f0a41d4d ed273297
......@@ -2,32 +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="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"
"""Deprecated legacy decision LLM module.
The active tutor orchestrator no longer calls this module. Keep it only for
manual legacy tests until it can be removed.
"""
import warnings
from app.deterministic_services import llm_client
warnings.warn(
"app.LLM_services.decision_LLM is deprecated and no longer used by the tutor orchestrator.",
DeprecationWarning,
stacklevel=2,
)
def context_decision(needs_more_context: bool, reason: str) -> dict:
# Kannst auch einfach nur return {"needs_more_context": needs_more_context, "reason": reason}
......
"""Deprecated legacy math-intent LLM module.
The active tutor orchestrator no longer calls this module. Keep it only for
manual legacy tests until it can be removed.
"""
import warnings
import sympy as sp
from app.deterministic_services import llm_client
warnings.warn(
"app.LLM_services.math_intent_LLM is deprecated and no longer used by the tutor orchestrator.",
DeprecationWarning,
stacklevel=2,
)
def sympy_solve(task: str, input: str, symbols: list[str] | None = None) -> str:
"""
......
......@@ -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,6 +7,9 @@ from typing import Optional
load_dotenv()
SUPPORTED_LLM_PROVIDERS = {"openai", "gwdg", "mistral", "ollama"}
SUPPORTED_EMBEDDING_PROVIDERS = {"sentence-transformer", "openai", "gwdg"}
class EmbeddingSettings(BaseModel):
embedding_type: str # "openai-like" oder "sentence-transformer"
......@@ -27,31 +30,83 @@ def get_retrieval_impl() -> str:
return "child"
def get_llm_provider() -> str:
value = os.getenv("LLM_PROVIDER")
if value is None or not value.strip():
supported = ", ".join(sorted(SUPPORTED_LLM_PROVIDERS))
raise ValueError(f"Missing LLM_PROVIDER. Expected one of: {supported}")
provider = value.strip().lower()
if provider not in SUPPORTED_LLM_PROVIDERS:
supported = ", ".join(sorted(SUPPORTED_LLM_PROVIDERS))
raise ValueError(
f"Unsupported LLM_PROVIDER: {value}. Expected one of: {supported}"
)
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)
......@@ -151,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:
......@@ -162,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")
......@@ -177,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.
......
import inspect
import json
import warnings
from datetime import date
from typing import Any, Callable
......@@ -58,11 +59,84 @@ def _record_call(result: dict, tokens: int | None = None) -> dict:
llm_quota.record_usage(pg_url, date.today(), calls=1, tokens=token_count)
return result
def _chat_openai(messages: list[dict]) -> dict:
def _warn_deprecated_provider_flags(use_ollama: bool, use_mistral: bool) -> None:
if use_ollama or use_mistral:
warnings.warn(
"use_ollama and use_mistral are deprecated and ignored. "
"Set LLM_PROVIDER in the environment instead.",
DeprecationWarning,
stacklevel=3,
)
def _warn_deprecated_tools() -> None:
warnings.warn(
"LLM toolcalling via llm_client.chat(..., tools=...) is deprecated.",
DeprecationWarning,
stacklevel=3,
)
def _ensure_within_llm_quota() -> None:
quota_settings = config.get_llm_quota_settings()
llm_quota.ensure_within_limits(
config.get_postgres_url(),
date.today(),
call_limit=quota_settings.daily_call_limit,
token_limit=quota_settings.daily_token_limit,
add_calls=1,
add_tokens=0,
)
def _quota_tracked_chat(chat_func: Callable[[], dict]) -> dict:
_ensure_within_llm_quota()
try:
result = chat_func()
except Exception:
_record_call({"raw": None}, tokens=0)
raise
if result:
return _record_call(result)
_record_call({"raw": None}, tokens=0)
return result
def _require_openai_chat_settings() -> config.OpenAIChatSettings:
settings = config.get_openai_chat_settings()
if not settings:
return {}
raise ValueError(
"LLM_PROVIDER=openai requires OPENAI_CHAT_MODEL, OPENAI_BASE_URL, "
"and OPENAI_API_KEY"
)
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:
raise ValueError(
"LLM_PROVIDER=mistral requires MISTRAL_CHAT_MODEL and MISTRAL_API_KEY"
)
return settings
def _chat_openai_compatible(
messages: list[dict],
settings: config.OpenAIChatSettings | None = None,
) -> dict:
settings = settings or _require_openai_chat_settings()
timeout = settings.timeout or 60.0
client = OpenAI(api_key=settings.api_key,
base_url=settings.base_url, timeout=timeout)
......@@ -74,11 +148,11 @@ def _chat_openai(messages: list[dict]) -> dict:
return {"raw": response, "message": message}
def _chat_mistral(messages: list[dict]) -> dict:
settings = config.get_mistral_chat_settings()
if not settings:
return {}
def _chat_mistral(
messages: list[dict],
settings: config.MistralChatSettings | None = None,
) -> dict:
settings = settings or _require_mistral_chat_settings()
kwargs: dict[str, Any] = {
"model": settings.model,
"messages": messages,
......@@ -97,89 +171,10 @@ def _chat_mistral(messages: list[dict]) -> dict:
return {"raw": response, "message": message}
def get_message_content(result: dict | object) -> str:
message = result.get("message") if isinstance(result, dict) else result
if isinstance(message, dict):
return message.get("content", "") or ""
if hasattr(message, "content"):
return getattr(message, "content") or ""
return ""
def chat(
def _chat_ollama(
messages: list[dict],
tools: list[Callable[..., Any]] | None = None,
use_ollama: bool = False,
use_mistral: bool = False,
) -> dict:
quota_settings = config.get_llm_quota_settings()
pg_url = config.get_postgres_url()
if use_mistral and tools:
raise RuntimeError("Mistral chat is currently only implemented for calls without tools.")
if use_mistral:
mistral_settings = config.get_mistral_chat_settings()
if not mistral_settings:
return {}
llm_quota.ensure_within_limits(
pg_url,
date.today(),
call_limit=quota_settings.daily_call_limit,
token_limit=quota_settings.daily_token_limit,
add_calls=1,
add_tokens=0,
)
try:
mistral_result = _chat_mistral(messages)
except Exception:
_record_call({"raw": None}, tokens=0)
raise
if mistral_result:
return _record_call(mistral_result)
_record_call({"raw": None}, tokens=0)
return mistral_result
if not tools and not use_ollama:
openai_settings = config.get_openai_chat_settings()
if openai_settings:
llm_quota.ensure_within_limits(
pg_url,
date.today(),
call_limit=quota_settings.daily_call_limit,
token_limit=quota_settings.daily_token_limit,
add_calls=1,
add_tokens=0,
)
try:
openai_result = _chat_openai(messages)
except Exception:
_record_call({"raw": None}, tokens=0)
raise
if openai_result:
return _record_call(openai_result)
_record_call({"raw": None}, tokens=0)
return openai_result
mistral_settings = config.get_mistral_chat_settings()
if mistral_settings:
llm_quota.ensure_within_limits(
pg_url,
date.today(),
call_limit=quota_settings.daily_call_limit,
token_limit=quota_settings.daily_token_limit,
add_calls=1,
add_tokens=0,
)
try:
mistral_result = _chat_mistral(messages)
except Exception:
_record_call({"raw": None}, tokens=0)
raise
if mistral_result:
return _record_call(mistral_result)
_record_call({"raw": None}, tokens=0)
return mistral_result
settings = config.get_ollama_settings()
client = ollama.Client(host=settings.base_url, timeout=settings.timeout)
......@@ -202,6 +197,46 @@ def chat(
return {"raw": response, "message": _extract_message(response)}
def get_message_content(result: dict | object) -> str:
message = result.get("message") if isinstance(result, dict) else result
if isinstance(message, dict):
return message.get("content", "") or ""
if hasattr(message, "content"):
return getattr(message, "content") or ""
return ""
def chat(
messages: list[dict],
tools: list[Callable[..., Any]] | None = None,
use_ollama: bool = False,
use_mistral: bool = False,
) -> dict:
if tools:
_warn_deprecated_tools()
_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_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))
if provider == "ollama":
return _chat_ollama(messages, tools=tools)
raise ValueError(f"Unsupported LLM_PROVIDER: {provider}")
def chat_with_tools(
messages: list[dict],
tools: list[Callable[..., Any]],
......@@ -209,6 +244,12 @@ def chat_with_tools(
use_mistral: bool = False,
return_after_tools: bool = False,
) -> tuple[dict, list[dict[str, Any]]]:
"""Deprecated legacy wrapper around model tool calls."""
warnings.warn(
"llm_client.chat_with_tools(...) is deprecated.",
DeprecationWarning,
stacklevel=2,
)
tool_map = {tool.__name__: tool for tool in tools}
result = chat(
messages=messages,
......@@ -235,6 +276,7 @@ def _apply_tool_calls(
messages: list[dict],
tool_map: dict[str, Callable[..., Any]],
) -> list[dict[str, Any]]:
"""Deprecated legacy helper for chat_with_tools."""
message = result.get("message") if isinstance(result, dict) else result
tool_calls = _extract_tool_calls(message)
outputs: list[dict[str, Any]] = []
......@@ -262,6 +304,7 @@ def _apply_tool_calls(
def _extract_tool_calls(message: object) -> list:
"""Deprecated legacy helper for chat_with_tools."""
if isinstance(message, dict):
return message.get("tool_calls") or []
if hasattr(message, "tool_calls"):
......@@ -270,6 +313,7 @@ def _extract_tool_calls(message: object) -> list:
def _tool_call_name_args(call: object) -> tuple[str, dict[str, Any]]:
"""Deprecated legacy helper for chat_with_tools."""
if isinstance(call, dict):
function = call.get("function") or {}
name = function.get("name") or ""
......@@ -282,6 +326,7 @@ def _tool_call_name_args(call: object) -> tuple[str, dict[str, Any]]:
def _parse_tool_arguments(arguments: object) -> dict[str, Any]:
"""Deprecated legacy helper for chat_with_tools."""
if isinstance(arguments, dict):
return arguments
if isinstance(arguments, str) and arguments.strip():
......
from __future__ import annotations
from app.LLM_services import decision_LLM, open_hint_LLM, math_intent_LLM, solver_LLM
from app.LLM_services import open_hint_LLM, solver_LLM
from app.deterministic_services import context_store
from app.deterministic_services.orchestrators import orchestrator_base as base
......@@ -8,21 +8,12 @@ from app.deterministic_services.orchestrators import orchestrator_base as base
def _on_bootstrap(state: base.ChatState, query_text: str) -> None:
base.bootstrap_retrieval(state.sheet, query_text, state.tool_log)
math_solution = base.log_timed_call(
state.tool_log,
"math_intent_LLM",
{"query": query_text},
lambda: math_intent_LLM.solve_with_tools(query_text),
)
if math_solution:
context_store.add_math_solution(state.sheet, math_solution)
# Sonstiges bei jedem Aufruf
def _on_turn_logic(state: base.ChatState) -> None:
sheet_text = context_store.format_sheet(state.sheet)
if state.new_chat:
sheet_text = context_store.format_sheet(state.sheet)
llm_solution = base.log_timed_call(
state.tool_log,
"LLM_Solution",
......@@ -30,19 +21,6 @@ def _on_turn_logic(state: base.ChatState) -> None:
lambda: solver_LLM.solve_question(state.last_user, sheet_text),
)
context_store.add_LLM_solution(state.sheet, llm_solution)
return
decision = base.log_timed_call(
state.tool_log,
"decision",
{"sheet": sheet_text},
lambda: decision_LLM.needs_more_context(sheet_text),
)
context_store.add_decision(state.sheet, decision)
if decision.get("needs_more_context"):
full_query = "\n".join(base.extract_user_messages(state.messages))
_on_bootstrap(state, full_query)
# Antwort generieren
......
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
......
import importlib.util
import os
import sys
import types
import unittest
import warnings
from unittest.mock import patch
os.environ.setdefault("EMBEDDING_TYPE", "sentence-transformer")
os.environ.setdefault("EMBEDDING_DIM", "512")
psycopg_stub = types.ModuleType("psycopg")
psycopg_rows_stub = types.ModuleType("psycopg.rows")
psycopg_rows_stub.dict_row = object()
psycopg_stub.rows = psycopg_rows_stub
pgvector_stub = types.ModuleType("pgvector")
pgvector_psycopg_stub = types.ModuleType("pgvector.psycopg")
pgvector_stub.Vector = list
pgvector_psycopg_stub.register_vector = lambda conn: None
ollama_stub = types.ModuleType("ollama")
ollama_stub.Client = object
openai_stub = types.ModuleType("openai")
openai_stub.OpenAI = object
mistralai_stub = types.ModuleType("mistralai")
mistralai_client_stub = types.ModuleType("mistralai.client")
mistralai_client_stub.Mistral = object
sentence_transformers_stub = types.ModuleType("sentence_transformers")
class _SentenceTransformerStub:
def __init__(self, *args, **kwargs) -> None:
self.max_seq_length = None
def encode(self, *args, **kwargs) -> list[list[float]]:
return [[0.0]]
sentence_transformers_stub.SentenceTransformer = _SentenceTransformerStub
if importlib.util.find_spec("psycopg") is None:
sys.modules.setdefault("psycopg", psycopg_stub)
sys.modules.setdefault("psycopg.rows", psycopg_rows_stub)
if importlib.util.find_spec("pgvector") is None:
sys.modules.setdefault("pgvector", pgvector_stub)
sys.modules.setdefault("pgvector.psycopg", pgvector_psycopg_stub)
if importlib.util.find_spec("ollama") is None:
sys.modules.setdefault("ollama", ollama_stub)
if importlib.util.find_spec("openai") is None:
sys.modules.setdefault("openai", openai_stub)
if importlib.util.find_spec("mistralai") is None:
sys.modules.setdefault("mistralai", mistralai_stub)
sys.modules.setdefault("mistralai.client", mistralai_client_stub)
else:
try:
import mistralai.client as installed_mistralai_client
except ImportError:
sys.modules.setdefault("mistralai.client", mistralai_client_stub)
else:
if not hasattr(installed_mistralai_client, "Mistral"):
installed_mistralai_client.Mistral = object
if importlib.util.find_spec("sentence_transformers") is None:
sys.modules.setdefault("sentence_transformers", sentence_transformers_stub)
from app import config
from app.deterministic_services import llm_client
from app.deterministic_services.orchestrators import orchestrator_tutor
from app.deterministic_services.orchestrators.orchestrator_base import ChatState
MESSAGES = [{"role": "user", "content": "Hallo"}]
def _dummy_tool() -> str:
return "ok"
class LLMProviderConfigTest(unittest.TestCase):
def test_get_llm_provider_accepts_supported_values(self) -> None:
for provider in ("openai", "gwdg", "mistral", "ollama"):
with self.subTest(provider=provider), patch.dict(
os.environ, {"LLM_PROVIDER": provider}, clear=True
):
self.assertEqual(config.get_llm_provider(), provider)
def test_get_llm_provider_normalizes_case_and_space(self) -> None:
with patch.dict(os.environ, {"LLM_PROVIDER": " OpenAI "}, clear=True):
self.assertEqual(config.get_llm_provider(), "openai")
def test_get_llm_provider_rejects_missing_value(self) -> None:
with patch.dict(os.environ, {}, clear=True):
with self.assertRaisesRegex(ValueError, "Missing LLM_PROVIDER"):
config.get_llm_provider()
def test_get_llm_provider_rejects_unknown_value(self) -> None:
with patch.dict(os.environ, {"LLM_PROVIDER": "anthropic"}, clear=True):
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:
settings = object()
expected = {"raw": object(), "message": {"content": "openai"}}
with patch.dict(os.environ, {"LLM_PROVIDER": "openai"}), patch.object(
llm_client, "_require_openai_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 openai_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)
openai_chat.assert_called_once_with(MESSAGES, settings)
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"}}
with patch.dict(os.environ, {"LLM_PROVIDER": "mistral"}), patch.object(
llm_client, "_require_mistral_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"
) as openai_chat, patch.object(
llm_client, "_chat_mistral", return_value=expected
) as mistral_chat, patch.object(
llm_client, "_chat_ollama"
) as ollama_chat:
result = llm_client.chat(MESSAGES)
self.assertEqual(result, expected)
openai_chat.assert_not_called()
mistral_chat.assert_called_once_with(MESSAGES, settings)
ollama_chat.assert_not_called()
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_compatible"
) as openai_chat, patch.object(
llm_client, "_chat_mistral"
) as mistral_chat, patch.object(
llm_client, "_chat_ollama", return_value=expected
) as ollama_chat:
result = llm_client.chat(MESSAGES)
self.assertEqual(result, expected)
openai_chat.assert_not_called()
mistral_chat.assert_not_called()
ollama_chat.assert_called_once_with(MESSAGES, tools=None)
def test_chat_deprecated_provider_flags_are_ignored(self) -> None:
settings = object()
expected = {"raw": object(), "message": {"content": "openai"}}
with patch.dict(os.environ, {"LLM_PROVIDER": "openai"}), patch.object(
llm_client, "_require_openai_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 openai_chat, patch.object(
llm_client, "_chat_ollama"
) as ollama_chat:
with self.assertWarns(DeprecationWarning):
result = llm_client.chat(MESSAGES, use_ollama=True, use_mistral=True)
self.assertEqual(result, expected)
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"
) as ensure_quota:
with self.assertRaisesRegex(ValueError, "LLM_PROVIDER=openai requires"):
llm_client.chat(MESSAGES)
ensure_quota.assert_not_called()
def test_chat_without_provider_does_not_fallback(self) -> None:
with patch.dict(os.environ, {}, clear=True), patch.object(
llm_client, "_chat_openai_compatible"
) as openai_chat, patch.object(
llm_client, "_chat_mistral"
) as mistral_chat, patch.object(
llm_client, "_chat_ollama"
) as ollama_chat:
with self.assertRaisesRegex(ValueError, "Missing LLM_PROVIDER"):
llm_client.chat(MESSAGES)
openai_chat.assert_not_called()
mistral_chat.assert_not_called()
ollama_chat.assert_not_called()
def test_chat_tools_argument_is_deprecated(self) -> None:
with patch.dict(os.environ, {"LLM_PROVIDER": "openai"}):
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always")
with self.assertRaisesRegex(RuntimeError, "toolcalling"):
llm_client.chat(MESSAGES, tools=[_dummy_tool])
messages = [str(warning.message) for warning in caught]
self.assertTrue(any("tools" in message for message in messages))
def test_chat_with_tools_is_deprecated(self) -> None:
response = {"raw": object(), "message": {"content": "ok"}}
with patch.dict(os.environ, {"LLM_PROVIDER": "ollama"}), patch.object(
llm_client, "_chat_ollama", return_value=response
):
with warnings.catch_warnings(record=True) as caught:
warnings.simplefilter("always")
result, tool_outputs = llm_client.chat_with_tools(
MESSAGES[:], [_dummy_tool], return_after_tools=True
)
messages = [str(warning.message) for warning in caught]
self.assertEqual(result, response)
self.assertEqual(tool_outputs, [])
self.assertTrue(any("chat_with_tools" in message for message in messages))
self.assertTrue(any("tools" in message for message in messages))
class TutorOrchestratorLegacyModuleTest(unittest.TestCase):
def test_tutor_orchestrator_does_not_import_legacy_llm_modules(self) -> None:
self.assertFalse(hasattr(orchestrator_tutor, "decision_LLM"))
self.assertFalse(hasattr(orchestrator_tutor, "math_intent_LLM"))
def test_bootstrap_only_runs_retrieval(self) -> None:
state = ChatState(
messages=MESSAGES[:],
draft=None,
chat_id="chat-1",
new_chat=True,
sheet={},
tool_log=[],
last_user="Hallo",
)
with patch.object(orchestrator_tutor.base, "bootstrap_retrieval") as bootstrap:
orchestrator_tutor._on_bootstrap(state, "Hallo")
bootstrap.assert_called_once_with(state.sheet, "Hallo", state.tool_log)
self.assertEqual(state.tool_log, [])
def test_non_new_turn_does_not_run_decision_llm(self) -> None:
state = ChatState(
messages=MESSAGES[:] + [{"role": "assistant", "content": "Antwort"}],
draft=None,
chat_id="chat-1",
new_chat=False,
sheet={"history": []},
tool_log=[],
last_user="Hallo",
)
with patch.object(orchestrator_tutor.solver_LLM, "solve_question") as solver:
orchestrator_tutor._on_turn_logic(state)
solver.assert_not_called()
self.assertEqual(state.tool_log, [])
if __name__ == "__main__":
unittest.main()
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,
......
import { useRef } from "react";
import { useEffect, useRef } from "react";
import { EditPencil, Send, Upload } from "iconoir-react";
import { t } from "../../i18n";
......@@ -23,6 +23,21 @@ export default function MessageInput({
}: MessageInputProps) {
const canSend = value.trim().length > 0;
const uploadInputRef = useRef<HTMLInputElement | null>(null);
const textareaRef = useRef<HTMLTextAreaElement | null>(null);
const wasSendingRef = useRef(isSending);
const shouldRestoreFocusRef = useRef(false);
const sendButtonStartedFromTextareaRef = useRef(false);
useEffect(() => {
if (wasSendingRef.current && !isSending) {
if (shouldRestoreFocusRef.current) {
textareaRef.current?.focus({ preventScroll: true });
}
shouldRestoreFocusRef.current = false;
}
wasSendingRef.current = isSending;
}, [isSending]);
const handleUploadClick = () => {
uploadInputRef.current?.click();
......@@ -39,6 +54,11 @@ export default function MessageInput({
await onUploadSolution?.(file);
};
const handleSend = (restoreFocus = document.activeElement === textareaRef.current) => {
shouldRestoreFocusRef.current = restoreFocus;
onSend();
};
return (
<div className="composer">
<div className="composer-row">
......@@ -72,7 +92,14 @@ export default function MessageInput({
<button
className="btn primary"
type="button"
onClick={onSend}
onPointerDown={() => {
sendButtonStartedFromTextareaRef.current =
document.activeElement === textareaRef.current;
}}
onClick={() => {
handleSend(sendButtonStartedFromTextareaRef.current);
sendButtonStartedFromTextareaRef.current = false;
}}
disabled={!canSend || isSending}
aria-label={t("send")}
title={t("send")}
......@@ -89,6 +116,7 @@ export default function MessageInput({
</button>
</div>
<textarea
ref={textareaRef}
className="composer-input"
placeholder={t("typeQuestionOrLatex")}
rows={3}
......@@ -116,7 +144,7 @@ export default function MessageInput({
if (event.key === "Enter" && !event.shiftKey) {
event.preventDefault();
onSend();
handleSend(true);
}
}}
/>
......
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