feat: add conversational RollCalc assistant
This commit is contained in:
@@ -0,0 +1,316 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user