578 lines
21 KiB
Python
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,
|
|
)
|