Commit 927e6a37 authored by Kantz's avatar Kantz
Browse files

erste version des Toolcallings

parent 4ff58000
......@@ -2,4 +2,5 @@
.env
__pycache__/
drawings/
markdown/
\ No newline at end of file
markdown/
logs/
\ No newline at end of file
......@@ -2,19 +2,15 @@ from __future__ import annotations
from typing import List, Optional
import os
import logging
from fastapi import APIRouter, HTTPException
from dotenv import load_dotenv
from pydantic import BaseModel, Field
from app.services import chat_engine
from app.services import orchestrator_engine
router = APIRouter()
load_dotenv()
OPENAI_BASE_URL = os.getenv("OPENAI_BASE_URL")
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
logger = logging.getLogger(__name__)
class ChatMessage(BaseModel):
......@@ -37,24 +33,16 @@ def chat(request: ChatRequest) -> ChatResponse:
if not request.messages:
raise HTTPException(status_code=400, detail="messages required")
if not OPENAI_BASE_URL or not OPENAI_API_KEY:
last = request.messages[-1]
reply = f"Mock reply to: {last.text}"
return ChatResponse(reply=reply, sources=["doc:example"])
try:
context = chat_engine.retrieve_context(
pg_url=os.getenv("POSTGRES_URL") or "",
query_text=request.messages[-1].text,
)
messages_payload = chat_engine.build_messages(
[{"role": m.role, "content": m.text} for m in request.messages],
context,
result = orchestrator_engine.run_chat(
[{"role": m.role, "content": m.text} for m in request.messages]
)
reply = chat_engine.call_chat_api(messages_payload)
reply = result["reply"]
sources = result["sources"]
except ValueError as exc:
raise HTTPException(status_code=500, detail=str(exc)) from exc
except Exception as exc:
logger.exception("Chat request failed")
raise HTTPException(status_code=502, detail="chat provider failed") from exc
return ChatResponse(reply=reply, sources=["openai-compatible"])
return ChatResponse(reply=reply, sources=sources)
from __future__ import annotations
import logging
import os
from typing import Any, Dict
import httpx
import psycopg
from fastapi import APIRouter
router = APIRouter()
logger = logging.getLogger(__name__)
def _check_ollama() -> dict:
base_url = os.getenv("OLLAMA_URL", "http://localhost:11434").rstrip("/")
url = f"{base_url}/api/tags"
try:
with httpx.Client(timeout=5.0) as client:
response = client.get(url)
response.raise_for_status()
return {"status": "ok", "url": url}
except Exception as exc:
return {"status": "error", "url": url, "detail": str(exc)}
def _normalize_openai_models_url(base_url: str) -> str:
trimmed = base_url.rstrip("/")
if trimmed.endswith("/v1/embeddings"):
trimmed = trimmed[: -len("/embeddings")]
if trimmed.endswith("/v1"):
return f"{trimmed}/models"
return f"{trimmed}/v1/models"
def _check_openai() -> dict:
base_url = os.getenv("OPENAI_BASE_URL")
api_key = os.getenv("OPENAI_API_KEY")
if not base_url or not api_key:
return {"status": "missing_config"}
url = _normalize_openai_models_url(base_url)
try:
headers = {"Authorization": f"Bearer {api_key}"}
with httpx.Client(timeout=5.0) as client:
response = client.get(url, headers=headers)
if response.status_code in (401, 403):
return {"status": "unauthorized", "url": url}
response.raise_for_status()
return {"status": "ok", "url": url}
except Exception as exc:
return {"status": "error", "url": url, "detail": str(exc)}
def _check_postgres() -> dict:
pg_url = os.getenv("POSTGRES_URL")
if not pg_url:
return {"status": "missing_config"}
try:
with psycopg.connect(pg_url, connect_timeout=5) as conn:
with conn.cursor() as cur:
cur.execute("SELECT 1")
cur.fetchone()
return {"status": "ok"}
except Exception as exc:
return {"status": "error", "detail": str(exc)}
@router.get("/api/health")
def health() -> Dict[str, Any]:
services = {
"ollama": _check_ollama(),
"openai": _check_openai(),
"postgres": _check_postgres(),
}
overall = "ok"
if any(value["status"] in ("error", "unauthorized") for value in services.values()):
overall = "degraded"
return {"status": overall, "services": services}
def run_startup_checks() -> Dict[str, Any]:
result = health()
status = result.get("status")
if status == "ok":
logger.info("Startup health check OK")
else:
logger.warning("Startup health check degraded: %s", result)
return result
import os
from dataclasses import dataclass
from dotenv import load_dotenv
load_dotenv()
@dataclass(frozen=True)
class OllamaSettings:
base_url: str
model: str
timeout: float | None
keepalive: str | None
@dataclass(frozen=True)
class EmbeddingSettings:
base_url: str
api_key: str
model: str
target_dim: int
def _read_float(value: str | None) -> float | None:
if value is None or value == "":
return None
try:
return float(value)
except ValueError:
return None
def get_ollama_settings() -> OllamaSettings:
return OllamaSettings(
base_url=os.getenv("OLLAMA_URL", "http://localhost:11434"),
model=os.getenv("OLLAMA_MODEL", "qwen3"),
timeout=_read_float(os.getenv("OLLAMA_TIMEOUT")),
keepalive=os.getenv("OLLAMA_KEEPALIVE"),
)
def get_embedding_settings() -> EmbeddingSettings:
base_url = os.getenv("OPENAI_BASE_URL")
api_key = os.getenv("OPENAI_API_KEY")
model = os.getenv("OPENAI_EMBED_MODEL", "text-embedding-3-large")
if not base_url or not api_key:
raise ValueError("Missing OPENAI_BASE_URL or OPENAI_API_KEY")
return EmbeddingSettings(
base_url=base_url,
api_key=api_key,
model=model,
target_dim=1024,
)
def get_postgres_url() -> str:
pg_url = os.getenv("POSTGRES_URL")
if not pg_url:
raise ValueError("Missing POSTGRES_URL")
return pg_url
from contextlib import asynccontextmanager
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from app.api import canvas, chat, retrieval
from app.api import canvas, chat, health, retrieval
@asynccontextmanager
async def lifespan(_app: FastAPI):
health.run_startup_checks()
yield
app = FastAPI(title="Math Tutor API", version="0.1.0")
app = FastAPI(title="Math Tutor API", version="0.1.0", lifespan=lifespan)
app.add_middleware(
CORSMiddleware,
......@@ -15,8 +23,4 @@ app.add_middleware(
app.include_router(chat.router)
app.include_router(canvas.router)
app.include_router(retrieval.router)
@app.get("/api/health")
def health() -> dict:
return {"status": "ok"}
app.include_router(health.router)
import inspect
import ollama
from app import config
def _filter_kwargs(func, kwargs: dict) -> dict:
try:
signature = inspect.signature(func)
except (TypeError, ValueError):
return kwargs
return {key: value for key, value in kwargs.items() if key in signature.parameters}
def chat(messages: list[dict], tools: list[dict] | None = None) -> dict:
settings = config.get_ollama_settings()
client = ollama.Client(host=settings.base_url, timeout=settings.timeout)
kwargs: dict = {"model": settings.model, "messages": messages}
if tools:
kwargs["tools"] = tools
if settings.keepalive:
kwargs["keep_alive"] = settings.keepalive
response = client.chat(**_filter_kwargs(client.chat, kwargs))
return {"raw": response, "message": _extract_message(response)}
def _extract_message(response) -> object:
if isinstance(response, dict) and "message" in response:
return response["message"]
if hasattr(response, "message"):
return response.message
return {}
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 get_tool_calls(result: dict | object) -> list:
message = result.get("message") if isinstance(result, dict) else result
if isinstance(message, dict):
return message.get("tool_calls") or []
if hasattr(message, "tool_calls"):
return getattr(message, "tool_calls") or []
return []
def normalize_tool_call(tool_call: object) -> dict:
if isinstance(tool_call, dict):
function = tool_call.get("function") or {}
return {"name": function.get("name"), "arguments": function.get("arguments")}
function = getattr(tool_call, "function", None)
if function:
return {"name": getattr(function, "name", None), "arguments": getattr(function, "arguments", None)}
return {"name": None, "arguments": None}
import json
from app.services import llm_client, retrieval_service, tool_registry, tool_logging
MAX_TOOL_STEPS = 4
def _normalize_args(raw_args):
if isinstance(raw_args, dict) and "arguments" in raw_args and isinstance(raw_args["arguments"], dict):
return raw_args["arguments"]
if isinstance(raw_args, str):
try:
return json.loads(raw_args)
except json.JSONDecodeError:
return {}
if isinstance(raw_args, dict):
return raw_args
return {}
def _system_messages(context: str) -> list[dict]:
return [
{"role": "system", "content": retrieval_service.SYSTEM_PROMPT},
{"role": "system", "content": context},
{
"role": "system",
"content": (
"Nutze Tools, wenn sie relevant sind. "
"Wenn eine mathematische Aufgabe enthalten ist, rufe sympy_solve auf. "
"Wenn zusaetzlicher Kontext benoetigt wird, rufe retrieve_context auf. "
"Loese Aufgaben nicht manuell."
),
},
]
def run_chat(messages: list[dict]) -> dict:
if not messages:
raise ValueError("messages required")
last_user = next((m for m in reversed(messages) if m.get("role") == "user"), None)
if not last_user:
raise ValueError("last user message required")
context, sources = retrieval_service.retrieve_context(query_text=last_user.get("content", ""))
chat_messages = _system_messages(context) + messages
tool_log: list[dict] = []
reply = ""
for _ in range(MAX_TOOL_STEPS):
result = llm_client.chat(chat_messages, tools=tool_registry.TOOL_SPECS)
tool_calls = llm_client.get_tool_calls(result)
if not tool_calls:
reply = llm_client.get_message_content(result)
break
for tool_call in tool_calls:
info = llm_client.normalize_tool_call(tool_call)
name = info.get("name")
args = _normalize_args(info.get("arguments"))
handler = tool_registry.TOOL_HANDLERS.get(name)
if not handler:
continue
tool_result = handler(**args)
tool_log.append({"name": name, "arguments": args, "response": tool_result})
if isinstance(tool_result, dict) and name == "retrieve_context":
context = tool_result.get("context", "")
sources = tool_result.get("sources", sources)
tool_content = context
else:
tool_content = tool_result
if not isinstance(tool_content, str):
tool_content = json.dumps(tool_content, ensure_ascii=True)
chat_messages.append({"role": "tool", "name": name, "content": tool_content})
if not reply:
reply = "Dazu steht nichts im Material"
tool_logging.write_tool_log(tool_log)
return {"reply": reply, "sources": sources, "tool_log": tool_log}
from __future__ import annotations
from typing import List
from app import config
from app.services.embeddings import OpenAILikeEmbeddings
from app.services import vector_store
SYSTEM_PROMPT = (
"Du bist ein Mathe-Tutor. Antworte auf Deutsch, klar und korrekt. "
"Nutze ausschliesslich den bereitgestellten Kontext. Wenn nichts zur Frage im Kontext steht, "
'antworte mit "Dazu steht nichts im Material" und nichts weiter. '
"Gib wenn moeglich eine kurze Struktur: (1) Idee, "
"(2) Definition, "
"(3) kurzer Begruendungs-/Rechenweg, "
"(4) Mini-Beispiel. "
"Zitiere Quellen inline mit den eckigen Klammern, die im Kontext vorangestellt sind, "
"z.B. [s2/ss1/c3 | definition | ...]."
)
CONTEXT_LIMITS = {
"direct": 8,
"indirect": 6,
"subsection": 2,
"section": 1,
}
def _get_embedder() -> OpenAILikeEmbeddings:
settings = config.get_embedding_settings()
return OpenAILikeEmbeddings(
base_url=settings.base_url,
api_key=settings.api_key,
model=settings.model,
target_dim=settings.target_dim,
)
def _format_ref(doc: vector_store.Retrieved) -> str:
meta = doc.metadata
sec = meta.get("section_index")
sub = meta.get("subsection_index")
child = meta.get("child_index")
ref = []
if sec is not None:
ref.append(f"s{sec}")
if sub is not None:
ref.append(f"ss{sub}")
if child is not None:
ref.append(f"c{child}")
ref_id = "/".join(ref) if ref else "unknown"
doc_type = meta.get("type") or doc.doc_type
title = (
meta.get("title")
or meta.get("subsection_title")
or meta.get("section_title")
or meta.get("path")
or "Untitled"
)
return f"[{ref_id} | {doc_type} | {title}]"
def _append_group(
label: str,
items: List[vector_store.Retrieved],
limit: int,
blocks: List[str],
sources: List[str],
) -> None:
if not items:
return
for doc in items[:limit]:
ref = _format_ref(doc)
blocks.append(f"{label} {ref}\n{doc.markdown}")
sources.append(ref)
def build_context(result: dict) -> tuple[str, List[str]]:
blocks: List[str] = []
sources: List[str] = []
_append_group("DIRECT", result.get("children_direct", []), CONTEXT_LIMITS["direct"], blocks, sources)
_append_group("INDIRECT", result.get("children_expanded", []), CONTEXT_LIMITS["indirect"], blocks, sources)
_append_group("SUBSECTION", result.get("subsections", []), CONTEXT_LIMITS["subsection"], blocks, sources)
_append_group("SECTION", result.get("sections", []), CONTEXT_LIMITS["section"], blocks, sources)
if not blocks:
return "KONTEXT: (leer)", sources
return "KONTEXT:\n" + "\n\n".join(blocks), sources
def retrieve_context(query_text: str, pg_url: str | None = None) -> tuple[str, List[str]]:
url = pg_url or config.get_postgres_url()
embedder = _get_embedder()
retrieval = vector_store.retrieve(
pg_url=url,
embedder=embedder,
query=query_text,
k=8,
expand_links=True,
)
return build_context(retrieval)
import json
import os
from datetime import datetime
def write_tool_log(entries: list[dict]) -> str:
os.makedirs("logs", exist_ok=True)
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
path = os.path.join("logs", f"tool_calls_{timestamp}.json")
payload = {
"timestamp": timestamp,
"entries": entries,
}
with open(path, "w", encoding="utf-8") as f:
json.dump(payload, f, ensure_ascii=True, indent=2)
return path
from app.tools import hint_tool, math_tool, retrieval_tool
TOOL_SPECS = [
math_tool.TOOL_SPEC,
hint_tool.TOOL_SPEC,
retrieval_tool.TOOL_SPEC,
]
TOOL_HANDLERS = {
"sympy_solve": math_tool.sympy_solve,
"generate_hint": hint_tool.generate_hint,
"retrieve_context": retrieval_tool.retrieve_context,
}
__all__ = [
"hint_tool",
"math_tool",
"retrieval_tool",
]
from app.services import llm_client
def generate_hint(task: str, solution: str, history: str | None = None) -> str:
prompt = (
"Du bist ein didaktischer Tutor. "
"Gib einen naechsten hilfreichen Hinweis, aber keine komplette Loesung. "
"Halte dich kurz und klar.\n\n"
"Aufgabe:\n"
+ task
+ "\n\n"
"Loesung (vom Mathe-Tool):\n"
+ solution
+ "\n"
)
if history:
prompt += "\nHistorie:\n" + history + "\n"
result = llm_client.chat(
messages=[{"role": "user", "content": prompt}],
)
return llm_client.get_message_content(result)
TOOL_SPEC = {
"type": "function",
"function": {
"name": "generate_hint",
"description": (
"Gibt einen didaktisch wertvollen naechsten Hinweis "
"auf Basis der Aufgabe und der berechneten Loesung."
),
"parameters": {
"type": "object",
"properties": {
"task": {"type": "string", "description": "Die gegebene Aufgabe"},
"solution": {
"type": "string",
"description": "Loesung aus dem Mathe-Tool",
},
"history": {
"type": "string",
"description": "Optionaler Verlauf, kann leer sein",
},
},
"required": ["task", "solution"],
},
},
}
import re
import sympy as sp
from sympy.parsing.latex import parse_latex
from sympy.parsing.sympy_parser import (
parse_expr,
standard_transformations,
implicit_multiplication_application,
)
transformations = standard_transformations + (implicit_multiplication_application,)
LATEX_HINTS = re.compile(r"(\\[a-zA-Z]+)|(\$[^$]+\$)|(\^\{)|(_\{)")
def looks_like_latex(text: str) -> bool:
return bool(LATEX_HINTS.search(text))
def parse_input(expr_text: str, sympy_symbols: dict[str, sp.Symbol]) -> sp.Expr:
if looks_like_latex(expr_text):
return parse_latex(expr_text)
return parse_expr(expr_text, transformations=transformations, local_dict=sympy_symbols)
def sympy_solve(task: str, input: str, symbols: list[str] | None = None) -> str:
try:
sympy_symbols: dict[str, sp.Symbol] = {}
if symbols:
for s in symbols:
sympy_symbols[s] = sp.symbols(s)
if "=" in input:
lhs, rhs = input.split("=")
expr = sp.Eq(parse_input(lhs, sympy_symbols), parse_input(rhs, sympy_symbols))
else:
expr = parse_input(input, sympy_symbols)
if task == "solve":
result = sp.solve(expr, list(sympy_symbols.values()) if symbols else None)
elif task == "simplify":
result = sp.simplify(expr)
elif task == "diff":
result = sp.diff(expr, *sympy_symbols.values())
elif task == "integrate":
result = sp.integrate(expr, *sympy_symbols.values())
else:
return f"Unsupported task: {task}"
return str(result)
except Exception as exc:
return f"SymPy error: {str(exc)}"
TOOL_SPEC = {
"type": "function",
"function": {
"name": "sympy_solve",
"description": (
"Loese oder bearbeite mathematische Ausdruecke mit SymPy. "
"Nutze dieses Tool, wenn eine mathematische Formel oder Gleichung erscheint."
),
"parameters": {
"type": "object",
"properties": {
"task": {
"type": "string",
"enum": ["solve", "simplify", "diff", "integrate"],
"description": "Mathematische Operation",
},
"input": {
"type": "string",
"description": "Mathematischer Ausdruck oder Gleichung, z. B. x**2 - 4 = 0",
},
"symbols": {
"type": "array",
"items": {"type": "string"},
"description": 'Variablen, z. B. ["x"]',
},
},
"required": ["task", "input"],
},
},
}
from app.services import retrieval_service
def retrieve_context(query: str) -> dict:
context, sources = retrieval_service.retrieve_context(query_text=query)
return {"context": context, "sources": sources}
TOOL_SPEC = {
"type": "function",
"function": {
"name": "retrieve_context",
"description": (
"Laedt Kontext aus dem Material via Vektor-Retrieval. "
"Nutze dieses Tool, wenn zusaetzlicher Kontext fuer die Antwort noetig ist."
),
"parameters": {
"type": "object",
"properties": {
"query": {"type": "string", "description": "Suchanfrage fuer das Retrieval"},
},
"required": ["query"],
},
},
}
......@@ -4,6 +4,8 @@ python-dotenv
mpxpy
pillow
httpx
ollama
sympy
psycopg[binary]
pgvector
pyyaml
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