Files

578 lines
21 KiB
Python

"""Strict Ollama-backed natural-language interpretation for RollCalc."""
from __future__ import annotations
from dataclasses import asdict, dataclass
import json
import math
import os
import re
import socket
import time
from typing import Any, Callable
from urllib.error import HTTPError, URLError
from urllib.request import Request, urlopen
DEFAULT_OLLAMA_URL = "http://127.0.0.1:11434"
DEFAULT_OLLAMA_MODEL = "qwen3.5:35B-A3B"
DEFAULT_TIMEOUT_SECONDS = 45.0
DEFAULT_MAX_RESPONSE_BYTES = 1_000_000
MAX_HTTP_ERROR_BODY_BYTES = 4_096
MAX_HTTP_ERROR_DIAGNOSTIC_CHARS = 1_000
NEW_CALCULATION_FIELDS = {
"intent",
"article_number",
"article_name_hint",
"roll_length_m",
"width_m",
"core_type",
"core_diameter_mm",
"include_roll_weight",
}
CHANGE_FIELDS = {
"article_number",
"article_name_hint",
"roll_length_m",
"width_m",
"core_type",
"core_diameter_mm",
"include_roll_weight",
}
NLU_JSON_SCHEMA: dict[str, Any] = {
"oneOf": [
{
"type": "object",
"additionalProperties": False,
"properties": {
"intent": {"const": "new_calculation"},
"article_number": {
"anyOf": [
{"type": "null"},
{"type": "string", "minLength": 1},
]
},
"article_name_hint": {
"anyOf": [
{"type": "null"},
{"type": "string", "minLength": 1},
]
},
"roll_length_m": {
"type": ["number", "null"],
"exclusiveMinimum": 0,
},
"width_m": {
"type": ["number", "null"],
"exclusiveMinimum": 0,
},
"core_type": {
"anyOf": [
{"type": "null"},
{"type": "string", "minLength": 1},
]
},
"core_diameter_mm": {
"type": ["number", "null"],
"exclusiveMinimum": 0,
},
"include_roll_weight": {"type": "boolean"},
},
"required": sorted(NEW_CALCULATION_FIELDS),
},
{
"type": "object",
"additionalProperties": False,
"properties": {
"intent": {"const": "modify_calculation"},
"changes": {
"type": "object",
"additionalProperties": False,
"minProperties": 1,
"properties": {
"article_number": {
"type": "string",
"minLength": 1,
},
"article_name_hint": {
"type": "string",
"minLength": 1,
},
"roll_length_m": {
"type": "number",
"exclusiveMinimum": 0,
},
"width_m": {
"type": "number",
"exclusiveMinimum": 0,
},
"core_type": {
"type": "string",
"minLength": 1,
},
"core_diameter_mm": {
"type": "number",
"exclusiveMinimum": 0,
},
"include_roll_weight": {"type": "boolean"},
},
},
},
"required": ["intent", "changes"],
},
{
"type": "object",
"additionalProperties": False,
"properties": {"intent": {"const": "unsupported"}},
"required": ["intent"],
},
]
}
class NLUError(RuntimeError):
"""Base error for controlled NLU failures."""
class NLUValidationError(NLUError):
"""Raised when model output is not valid under the narrow NLU schema."""
class OllamaUnavailableError(NLUError):
"""Raised when the configured local Ollama service cannot be reached."""
class OllamaTimeoutError(NLUError):
"""Raised when Ollama does not answer within the configured timeout."""
class OllamaResponseError(NLUError):
"""Raised when Ollama returns an unsuccessful or malformed response."""
def __init__(
self,
message: str,
*,
status_code: int | None = None,
) -> None:
super().__init__(message)
self.status_code = status_code
@dataclass(frozen=True)
class OllamaConfig:
base_url: str = DEFAULT_OLLAMA_URL
model: str = DEFAULT_OLLAMA_MODEL
timeout_seconds: float = DEFAULT_TIMEOUT_SECONDS
temperature: float = 0.0
think: bool = False
max_response_bytes: int = DEFAULT_MAX_RESPONSE_BYTES
@classmethod
def from_env(cls) -> "OllamaConfig":
return cls(
base_url=os.getenv("ROLLCALC_OLLAMA_URL", DEFAULT_OLLAMA_URL),
model=os.getenv("ROLLCALC_OLLAMA_MODEL", DEFAULT_OLLAMA_MODEL),
timeout_seconds=float(
os.getenv(
"ROLLCALC_OLLAMA_TIMEOUT_SECONDS",
str(DEFAULT_TIMEOUT_SECONDS),
)
),
temperature=float(os.getenv("ROLLCALC_OLLAMA_TEMPERATURE", "0")),
)
def __post_init__(self) -> None:
if not self.base_url.strip():
raise ValueError("Ollama base URL is required")
if not self.model.strip():
raise ValueError("Ollama model is required")
if not math.isfinite(self.timeout_seconds) or self.timeout_seconds <= 0:
raise ValueError("Ollama timeout must be greater than zero")
if not math.isfinite(self.temperature) or self.temperature < 0:
raise ValueError("Ollama temperature must be zero or greater")
if self.max_response_bytes <= 0:
raise ValueError("Ollama response limit must be greater than zero")
@dataclass(frozen=True)
class NLUInterpretation:
intent: str
article_number: str | None = None
article_name_hint: str | None = None
roll_length_m: float | None = None
width_m: float | None = None
core_type: str | None = None
core_diameter_mm: float | None = None
include_roll_weight: bool = False
changes: dict[str, Any] | None = None
def to_dict(self) -> dict[str, Any]:
if self.intent == "modify_calculation":
return {"intent": self.intent, "changes": dict(self.changes or {})}
if self.intent == "unsupported":
return {"intent": self.intent}
result = asdict(self)
result.pop("changes")
return result
@dataclass(frozen=True)
class NLUResponse:
interpretation: NLUInterpretation
raw_model_json: str
latency_ms: float
model: str
@dataclass(frozen=True)
class ClarificationInterpretation:
status: str
field: str
value: Any = None
def to_dict(self) -> dict[str, Any]:
result = {"status": self.status, "field": self.field}
if self.status == "parsed":
result["value"] = self.value
return result
_CLARIFICATION_UNITS = {
"roll_length_m": {"", "m", "meter", "metern"},
"width_m": {"", "m", "meter", "metern"},
"thickness_mm": {"", "mm", "millimeter", "millimetern"},
"core_diameter_mm": {"", "mm", "millimeter", "millimetern"},
"area_weight_g_m2": {"", "g/m2", "gsm"},
}
_CLARIFICATION_TEXT_FIELDS = {"article_number", "article_name_hint", "core_type"}
_CLARIFICATION_NUMBER = re.compile(
r"^\s*(?P<number>\d+(?:[.,]\d+)?)\s*"
r"(?P<unit>[A-Za-zÀ-ÖØ-öø-ÿ²/^0-9]+)?\s*[.!]?\s*$"
)
def parse_clarification_reply(
text: str,
expected_field: str,
) -> ClarificationInterpretation:
"""Parse one reply only as the application-selected clarification field."""
supported_fields = set(_CLARIFICATION_UNITS) | _CLARIFICATION_TEXT_FIELDS
if expected_field not in supported_fields:
raise NLUValidationError("unsupported clarification field")
if not isinstance(text, str) or not text.strip():
return ClarificationInterpretation("incompatible", expected_field)
if expected_field in _CLARIFICATION_TEXT_FIELDS:
try:
value = _required_text(text, expected_field)
except NLUValidationError:
return ClarificationInterpretation("incompatible", expected_field)
return ClarificationInterpretation("parsed", expected_field, value)
match = _CLARIFICATION_NUMBER.fullmatch(text)
if match is None:
return ClarificationInterpretation("incompatible", expected_field)
unit = (match.group("unit") or "").casefold()
unit = unit.replace("²", "2").replace("^", "")
if unit not in _CLARIFICATION_UNITS[expected_field]:
return ClarificationInterpretation("incompatible", expected_field)
value = float(match.group("number").replace(",", "."))
try:
value = _positive_number(value, expected_field, optional=False)
except NLUValidationError:
return ClarificationInterpretation("incompatible", expected_field)
return ClarificationInterpretation("parsed", expected_field, value)
def _optional_text(value: Any, field: str) -> str | None:
if value is None:
return None
if not isinstance(value, str):
raise NLUValidationError(f"{field} must be text or null")
value = " ".join(value.split()).strip()
if not value:
return None
if len(value) > 240:
raise NLUValidationError(f"{field} is too long")
return value
def _required_text(value: Any, field: str) -> str:
if not isinstance(value, str):
raise NLUValidationError(f"{field} must be text")
value = " ".join(value.split()).strip()
if not value:
raise NLUValidationError(f"{field} must not be empty")
if len(value) > 240:
raise NLUValidationError(f"{field} is too long")
return value
def _positive_number(value: Any, field: str, *, optional: bool) -> float | None:
if value is None and optional:
return None
if isinstance(value, bool) or not isinstance(value, (int, float)):
suffix = " or null" if optional else ""
raise NLUValidationError(f"{field} must be a number{suffix}")
value = float(value)
if not math.isfinite(value) or value <= 0:
raise NLUValidationError(f"{field} must be greater than zero")
return value
def validate_nlu_payload(payload: Any) -> NLUInterpretation:
"""Validate untrusted model output without coercing or inferring values."""
if not isinstance(payload, dict):
raise NLUValidationError("model output must be a JSON object")
intent = payload.get("intent")
if intent not in {"new_calculation", "modify_calculation", "unsupported"}:
raise NLUValidationError("unsupported or missing NLU intent")
if intent == "unsupported":
unknown = sorted(set(payload) - {"intent"})
if unknown:
raise NLUValidationError(
"unknown unsupported-intent fields: " + ", ".join(unknown)
)
return NLUInterpretation(intent="unsupported")
if intent == "new_calculation":
unknown = sorted(set(payload) - NEW_CALCULATION_FIELDS)
missing = sorted(NEW_CALCULATION_FIELDS - set(payload))
if unknown:
raise NLUValidationError(
"unknown new-calculation fields: " + ", ".join(unknown)
)
if missing:
raise NLUValidationError(
"missing new-calculation fields: " + ", ".join(missing)
)
if not isinstance(payload["include_roll_weight"], bool):
raise NLUValidationError("include_roll_weight must be boolean")
article_number = _optional_text(payload["article_number"], "article_number")
return NLUInterpretation(
intent=intent,
article_number=article_number,
article_name_hint=_optional_text(
payload["article_name_hint"], "article_name_hint"
),
roll_length_m=_positive_number(
payload["roll_length_m"], "roll_length_m", optional=True
),
width_m=_positive_number(payload["width_m"], "width_m", optional=True),
core_type=_optional_text(payload["core_type"], "core_type"),
core_diameter_mm=_positive_number(
payload["core_diameter_mm"], "core_diameter_mm", optional=True
),
include_roll_weight=payload["include_roll_weight"],
)
unknown = sorted(set(payload) - {"intent", "changes"})
if unknown:
raise NLUValidationError("unknown modification fields: " + ", ".join(unknown))
changes = payload.get("changes")
if not isinstance(changes, dict) or not changes:
raise NLUValidationError("changes must be a non-empty JSON object")
unknown_changes = sorted(set(changes) - CHANGE_FIELDS)
if unknown_changes:
raise NLUValidationError("unknown change fields: " + ", ".join(unknown_changes))
validated_changes: dict[str, Any] = {}
for field, value in changes.items():
if field in {"roll_length_m", "width_m", "core_diameter_mm"}:
validated_changes[field] = _positive_number(value, field, optional=False)
elif field == "include_roll_weight":
if not isinstance(value, bool):
raise NLUValidationError("include_roll_weight must be boolean")
validated_changes[field] = value
else:
validated_changes[field] = _required_text(value, field)
return NLUInterpretation(
intent="modify_calculation",
changes=validated_changes,
)
def parse_nlu_json(raw_model_json: str) -> NLUInterpretation:
if not isinstance(raw_model_json, str) or not raw_model_json.strip():
raise NLUValidationError("model output is empty")
try:
payload = json.loads(raw_model_json)
except json.JSONDecodeError as error:
raise NLUValidationError("model output is not valid JSON") from error
return validate_nlu_payload(payload)
def _http_error_diagnostic(error: HTTPError) -> str:
try:
raw_body = error.read(MAX_HTTP_ERROR_BODY_BYTES + 1)
except Exception:
return ""
if not raw_body:
return ""
if isinstance(raw_body, bytes):
body = raw_body.decode("utf-8", errors="replace")
else:
body = str(raw_body)
body = " ".join(body.split())
if not body:
return ""
diagnostic = body
try:
parsed_body = json.loads(body)
except json.JSONDecodeError:
pass
else:
if isinstance(parsed_body, dict):
ollama_error = parsed_body.get("error")
if isinstance(ollama_error, str) and ollama_error.strip():
diagnostic = " ".join(ollama_error.split())
if len(diagnostic) > MAX_HTTP_ERROR_DIAGNOSTIC_CHARS:
diagnostic = diagnostic[: MAX_HTTP_ERROR_DIAGNOSTIC_CHARS - 3] + "..."
return diagnostic
def _system_prompt(
*,
has_state: bool,
expected_fields: tuple[str, ...],
current_state: dict[str, Any] | None,
) -> str:
context = {
"has_calculation_state": has_state,
"expected_clarification_fields": list(expected_fields),
"current_calculation": current_state,
}
return (
"You are the strictly constrained NLU adapter for RollCalc. "
"Interpret the user's German or English text; never calculate, estimate, "
"or invent engineering values, product data, warnings, defaults, or core "
"mappings. Extract a user-provided product phrase only as article_name_hint; "
"never invent, rank, or select article candidates. Return only JSON matching "
"the supplied schema. Preserve article "
"numbers as strings. Convert an explicitly written German decimal comma to "
"a JSON number. For every optional textual field the user did not provide, "
"return JSON null; never return an empty or whitespace-only string. Set "
"include_roll_weight=true only when the user explicitly asks for roll weight. "
"Use modify_calculation for follow-ups that change existing structured state; "
"put every explicitly requested changed input into changes. "
"When current_calculation is present, references such as die Rolle, dieselbe "
"Rolle, sie, bei 80 m, mit 5 m Breite, wie schwer ist sie, or nimm einen "
"Stahlkern refer to that calculation and must not discard its other inputs. "
"Use new_calculation only when the user explicitly starts a new calculation "
"or identifies a different article. Do not copy unchanged context values into "
"the output. "
"If expected clarification fields are listed, interpret a short answer only "
"against those fields. Use unsupported for every other task. "
"Never emit diameter results, weight results, warnings, explanations, or "
"markdown. Technical context: "
+ json.dumps(context, ensure_ascii=False, separators=(",", ":"))
+ ". Required JSON schema: "
+ json.dumps(NLU_JSON_SCHEMA, ensure_ascii=False, separators=(",", ":"))
)
class OllamaNLUClient:
def __init__(
self,
config: OllamaConfig | None = None,
*,
opener: Callable[..., Any] = urlopen,
) -> None:
self.config = config or OllamaConfig.from_env()
self._opener = opener
def interpret(
self,
text: str,
*,
has_state: bool = False,
expected_fields: tuple[str, ...] = (),
current_state: dict[str, Any] | None = None,
) -> NLUResponse:
if not isinstance(text, str) or not text.strip():
raise NLUValidationError("message must be non-empty text")
if len(text) > 2_000:
raise NLUValidationError("message is too long")
payload = {
"model": self.config.model,
"messages": [
{
"role": "system",
"content": _system_prompt(
has_state=has_state,
expected_fields=expected_fields,
current_state=current_state,
),
},
{"role": "user", "content": text.strip()},
],
"stream": False,
"think": self.config.think,
"format": NLU_JSON_SCHEMA,
"options": {
"temperature": self.config.temperature,
"num_predict": 300,
},
}
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
request = Request(
self.config.base_url.rstrip("/") + "/api/chat",
data=body,
headers={"Content-Type": "application/json"},
method="POST",
)
started = time.monotonic()
try:
with self._opener(request, timeout=self.config.timeout_seconds) as response:
raw_response = response.read(self.config.max_response_bytes + 1)
except (socket.timeout, TimeoutError) as error:
raise OllamaTimeoutError("Ollama request timed out") from error
except HTTPError as error:
diagnostic = _http_error_diagnostic(error)
message = f"Ollama returned HTTP {error.code}"
if diagnostic:
message += f": {diagnostic}"
raise OllamaResponseError(
message,
status_code=error.code,
) from error
except URLError as error:
if isinstance(error.reason, (socket.timeout, TimeoutError)):
raise OllamaTimeoutError("Ollama request timed out") from error
raise OllamaUnavailableError("Ollama service is unavailable") from error
except OSError as error:
raise OllamaUnavailableError("Ollama service is unavailable") from error
latency_ms = (time.monotonic() - started) * 1_000
if len(raw_response) > self.config.max_response_bytes:
raise OllamaResponseError("Ollama response is too large")
try:
response_payload = json.loads(raw_response.decode("utf-8"))
except (UnicodeDecodeError, json.JSONDecodeError) as error:
raise OllamaResponseError("Ollama returned invalid JSON") from error
message = (
response_payload.get("message")
if isinstance(response_payload, dict)
else None
)
content = message.get("content") if isinstance(message, dict) else None
if not isinstance(content, str):
raise OllamaResponseError("Ollama response has no message content")
interpretation = parse_nlu_json(content)
return NLUResponse(
interpretation=interpretation,
raw_model_json=content,
latency_ms=latency_ms,
model=self.config.model,
)