317 lines
11 KiB
Python
317 lines
11 KiB
Python
import io
|
|
import json
|
|
import os
|
|
import socket
|
|
import unittest
|
|
from unittest.mock import patch
|
|
from urllib.error import HTTPError, URLError
|
|
|
|
from ollama_nlu import (
|
|
NLUValidationError,
|
|
OllamaConfig,
|
|
OllamaNLUClient,
|
|
OllamaResponseError,
|
|
OllamaTimeoutError,
|
|
OllamaUnavailableError,
|
|
parse_clarification_reply,
|
|
parse_nlu_json,
|
|
validate_nlu_payload,
|
|
)
|
|
|
|
|
|
def new_calculation_payload(**changes):
|
|
payload = {
|
|
"intent": "new_calculation",
|
|
"article_number": "180205",
|
|
"article_name_hint": "Bentofix NSP 4900",
|
|
"roll_length_m": 65.0,
|
|
"width_m": None,
|
|
"core_type": None,
|
|
"core_diameter_mm": None,
|
|
"include_roll_weight": False,
|
|
}
|
|
payload.update(changes)
|
|
return payload
|
|
|
|
|
|
class FakeResponse(io.BytesIO):
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc_value, traceback):
|
|
self.close()
|
|
|
|
|
|
class RecordingOpener:
|
|
def __init__(self, model_payload):
|
|
self.model_payload = model_payload
|
|
self.request_payload = None
|
|
self.timeout = None
|
|
|
|
def __call__(self, request, *, timeout):
|
|
self.request_payload = json.loads(request.data.decode("utf-8"))
|
|
self.timeout = timeout
|
|
response = {
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": json.dumps(self.model_payload),
|
|
},
|
|
"done": True,
|
|
}
|
|
return FakeResponse(json.dumps(response).encode("utf-8"))
|
|
|
|
|
|
class NLUValidationTests(unittest.TestCase):
|
|
def test_core_diameter_clarification_accepts_mm_and_unitless_values(self):
|
|
with_unit = parse_clarification_reply("150 mm", "core_diameter_mm")
|
|
unitless = parse_clarification_reply("150", "core_diameter_mm")
|
|
|
|
self.assertEqual(
|
|
with_unit.to_dict(),
|
|
{"status": "parsed", "field": "core_diameter_mm", "value": 150.0},
|
|
)
|
|
self.assertEqual(unitless.to_dict(), with_unit.to_dict())
|
|
|
|
def test_core_diameter_clarification_rejects_meter_unit(self):
|
|
interpretation = parse_clarification_reply("4,90 m", "core_diameter_mm")
|
|
|
|
self.assertEqual(
|
|
interpretation.to_dict(),
|
|
{"status": "incompatible", "field": "core_diameter_mm"},
|
|
)
|
|
|
|
def test_length_and_width_clarifications_use_expected_meter_field(self):
|
|
length = parse_clarification_reply("80 m", "roll_length_m")
|
|
width = parse_clarification_reply("4,90 m", "width_m")
|
|
|
|
self.assertEqual(length.value, 80.0)
|
|
self.assertEqual(length.field, "roll_length_m")
|
|
self.assertEqual(width.value, 4.9)
|
|
self.assertEqual(width.field, "width_m")
|
|
|
|
def test_empty_optional_article_name_is_normalized_to_none(self):
|
|
interpretation = validate_nlu_payload(
|
|
new_calculation_payload(article_name_hint="")
|
|
)
|
|
|
|
self.assertIsNone(interpretation.article_name_hint)
|
|
|
|
def test_whitespace_optional_article_name_is_normalized_to_none(self):
|
|
interpretation = validate_nlu_payload(
|
|
new_calculation_payload(article_name_hint=" ")
|
|
)
|
|
|
|
self.assertIsNone(interpretation.article_name_hint)
|
|
|
|
def test_null_optional_article_name_remains_none(self):
|
|
interpretation = validate_nlu_payload(
|
|
new_calculation_payload(article_name_hint=None)
|
|
)
|
|
|
|
self.assertIsNone(interpretation.article_name_hint)
|
|
|
|
def test_valid_optional_article_name_is_preserved(self):
|
|
interpretation = validate_nlu_payload(
|
|
new_calculation_payload(article_name_hint="Bentofix NSP 4900")
|
|
)
|
|
|
|
self.assertEqual(interpretation.article_name_hint, "Bentofix NSP 4900")
|
|
|
|
def test_empty_text_is_still_invalid_for_a_selected_modification(self):
|
|
with self.assertRaises(NLUValidationError):
|
|
validate_nlu_payload(
|
|
{
|
|
"intent": "modify_calculation",
|
|
"changes": {"core_type": " "},
|
|
}
|
|
)
|
|
|
|
def test_new_calculation_extracts_requested_german_fields(self):
|
|
interpretation = validate_nlu_payload(new_calculation_payload())
|
|
|
|
self.assertEqual(interpretation.intent, "new_calculation")
|
|
self.assertEqual(interpretation.article_number, "180205")
|
|
self.assertEqual(interpretation.article_name_hint, "Bentofix NSP 4900")
|
|
self.assertEqual(interpretation.roll_length_m, 65.0)
|
|
self.assertIsNone(interpretation.width_m)
|
|
self.assertIsNone(interpretation.core_diameter_mm)
|
|
|
|
def test_length_modification_contains_only_the_requested_change(self):
|
|
interpretation = validate_nlu_payload(
|
|
{
|
|
"intent": "modify_calculation",
|
|
"changes": {"roll_length_m": 80.0},
|
|
}
|
|
)
|
|
|
|
self.assertEqual(
|
|
interpretation.to_dict(),
|
|
{
|
|
"intent": "modify_calculation",
|
|
"changes": {"roll_length_m": 80.0},
|
|
},
|
|
)
|
|
|
|
def test_steel_core_modification_never_contains_a_guessed_diameter(self):
|
|
interpretation = validate_nlu_payload(
|
|
{
|
|
"intent": "modify_calculation",
|
|
"changes": {"core_type": "steel"},
|
|
}
|
|
)
|
|
|
|
self.assertEqual(interpretation.changes, {"core_type": "steel"})
|
|
self.assertNotIn("core_diameter_mm", interpretation.changes)
|
|
|
|
def test_modification_accepts_multiple_explicit_allowed_changes(self):
|
|
interpretation = validate_nlu_payload(
|
|
{
|
|
"intent": "modify_calculation",
|
|
"changes": {
|
|
"roll_length_m": 80.0,
|
|
"width_m": 4.9,
|
|
},
|
|
}
|
|
)
|
|
|
|
self.assertEqual(
|
|
interpretation.changes,
|
|
{"roll_length_m": 80.0, "width_m": 4.9},
|
|
)
|
|
|
|
def test_forbidden_calculated_fields_are_rejected_at_every_level(self):
|
|
with self.assertRaises(NLUValidationError):
|
|
validate_nlu_payload(
|
|
new_calculation_payload(average_diameter_mm=999.0)
|
|
)
|
|
with self.assertRaises(NLUValidationError):
|
|
validate_nlu_payload(
|
|
{
|
|
"intent": "modify_calculation",
|
|
"changes": {"roll_weight_kg": 1.0},
|
|
}
|
|
)
|
|
|
|
def test_invalid_json_is_handled_without_fallback_parsing(self):
|
|
with self.assertRaises(NLUValidationError):
|
|
parse_nlu_json("```json\n{}\n```")
|
|
|
|
|
|
class OllamaClientTests(unittest.TestCase):
|
|
def test_model_uses_available_default_and_environment_override(self):
|
|
self.assertEqual(OllamaConfig().model, "qwen3.5:35B-A3B")
|
|
with patch.dict(
|
|
os.environ,
|
|
{"ROLLCALC_OLLAMA_MODEL": "custom-qwen:model"},
|
|
):
|
|
self.assertEqual(
|
|
OllamaConfig.from_env().model,
|
|
"custom-qwen:model",
|
|
)
|
|
|
|
def test_german_request_uses_constrained_non_thinking_chat_request(self):
|
|
opener = RecordingOpener(new_calculation_payload())
|
|
config = OllamaConfig(
|
|
base_url="http://ollama.test:11434",
|
|
model="custom-qwen:model",
|
|
timeout_seconds=12,
|
|
)
|
|
client = OllamaNLUClient(config, opener=opener)
|
|
|
|
response = client.interpret(
|
|
"Welchen Durchmesser hat Bentofix NSP 4900, "
|
|
"Artikelnummer 180205 bei 65 m Länge?"
|
|
)
|
|
|
|
self.assertEqual(response.interpretation.article_number, "180205")
|
|
self.assertEqual(response.interpretation.roll_length_m, 65.0)
|
|
self.assertEqual(
|
|
opener.request_payload["messages"][1]["content"],
|
|
"Welchen Durchmesser hat Bentofix NSP 4900, "
|
|
"Artikelnummer 180205 bei 65 m Länge?",
|
|
)
|
|
self.assertEqual(opener.request_payload["model"], "custom-qwen:model")
|
|
self.assertEqual(opener.request_payload["options"]["temperature"], 0.0)
|
|
self.assertFalse(opener.request_payload["think"])
|
|
self.assertFalse(opener.request_payload["stream"])
|
|
self.assertIsInstance(opener.request_payload["format"], dict)
|
|
self.assertIn(
|
|
"never return an empty or whitespace-only string",
|
|
opener.request_payload["messages"][0]["content"],
|
|
)
|
|
self.assertEqual(opener.timeout, 12)
|
|
|
|
def test_generated_ollama_schema_has_no_regex_patterns(self):
|
|
opener = RecordingOpener(new_calculation_payload())
|
|
|
|
OllamaNLUClient(opener=opener).interpret("Artikel 180205 mit 65 Metern")
|
|
|
|
encoded_schema = json.dumps(opener.request_payload["format"])
|
|
self.assertNotIn('"pattern"', encoded_schema)
|
|
|
|
def test_current_input_state_is_provided_without_calculated_results(self):
|
|
opener = RecordingOpener(
|
|
{
|
|
"intent": "modify_calculation",
|
|
"changes": {"width_m": 5.0},
|
|
}
|
|
)
|
|
current_state = {
|
|
"has_successful_calculation": True,
|
|
"article_number": "180205",
|
|
"roll_length_m": 80.0,
|
|
"width_m": None,
|
|
"core_type": "194mm Stahl",
|
|
"core_diameter_mm": 194.0,
|
|
"include_roll_weight": False,
|
|
}
|
|
|
|
OllamaNLUClient(opener=opener).interpret(
|
|
"Wieviel wiegt die Rolle bei 5 m Breite?",
|
|
has_state=True,
|
|
current_state=current_state,
|
|
)
|
|
|
|
system_prompt = opener.request_payload["messages"][0]["content"]
|
|
self.assertIn('"current_calculation":{', system_prompt)
|
|
self.assertIn('"article_number":"180205"', system_prompt)
|
|
self.assertIn('"roll_length_m":80.0', system_prompt)
|
|
self.assertIn("die Rolle, dieselbe Rolle", system_prompt)
|
|
self.assertNotIn("average_diameter_mm", system_prompt)
|
|
self.assertNotIn("roll_weight_kg", system_prompt)
|
|
|
|
def test_timeout_and_unavailable_service_are_controlled(self):
|
|
def timeout_opener(request, *, timeout):
|
|
raise socket.timeout()
|
|
|
|
def unavailable_opener(request, *, timeout):
|
|
raise URLError(ConnectionRefusedError())
|
|
|
|
with self.assertRaises(OllamaTimeoutError):
|
|
OllamaNLUClient(opener=timeout_opener).interpret("Test")
|
|
with self.assertRaises(OllamaUnavailableError):
|
|
OllamaNLUClient(opener=unavailable_opener).interpret("Test")
|
|
|
|
def test_http_error_includes_status_and_ollama_response_detail(self):
|
|
def bad_request_opener(request, *, timeout):
|
|
raise HTTPError(
|
|
request.full_url,
|
|
400,
|
|
"Bad Request",
|
|
hdrs=None,
|
|
fp=io.BytesIO(
|
|
b'{"error":"some Ollama explanation"}'
|
|
),
|
|
)
|
|
|
|
with self.assertRaises(OllamaResponseError) as raised:
|
|
OllamaNLUClient(opener=bad_request_opener).interpret("Test")
|
|
|
|
self.assertEqual(raised.exception.status_code, 400)
|
|
self.assertIn("HTTP 400", str(raised.exception))
|
|
self.assertIn("some Ollama explanation", str(raised.exception))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|