113 lines
3.8 KiB
Python
113 lines
3.8 KiB
Python
import base64
|
|
import os
|
|
import tempfile
|
|
import unittest
|
|
|
|
import app as rollcalc_app
|
|
from conversation_service import ConversationService, InMemoryConversationStore
|
|
from ollama_nlu import NLUResponse, OllamaUnavailableError, validate_nlu_payload
|
|
|
|
|
|
class StaticNLUClient:
|
|
def __init__(self, payload):
|
|
self.payload = payload
|
|
|
|
def interpret(
|
|
self,
|
|
message,
|
|
*,
|
|
has_state=False,
|
|
expected_fields=(),
|
|
current_state=None,
|
|
):
|
|
if isinstance(self.payload, Exception):
|
|
raise self.payload
|
|
return NLUResponse(
|
|
interpretation=validate_nlu_payload(self.payload),
|
|
raw_model_json="raw-json",
|
|
latency_ms=1.0,
|
|
model="test-model",
|
|
)
|
|
|
|
|
|
class ConversationApiTests(unittest.TestCase):
|
|
def setUp(self):
|
|
self.temp_dir = tempfile.TemporaryDirectory()
|
|
self.original_log_file = rollcalc_app.LOG_FILE
|
|
self.original_service = rollcalc_app.CONVERSATION_SERVICE
|
|
rollcalc_app.LOG_FILE = os.path.join(self.temp_dir.name, "access.json")
|
|
rollcalc_app.app.config.update(TESTING=True)
|
|
self.client = rollcalc_app.app.test_client()
|
|
credentials = base64.b64encode(b"mtazl:rollcalc").decode("ascii")
|
|
self.headers = {"Authorization": f"Basic {credentials}"}
|
|
|
|
def tearDown(self):
|
|
rollcalc_app.LOG_FILE = self.original_log_file
|
|
rollcalc_app.CONVERSATION_SERVICE = self.original_service
|
|
self.temp_dir.cleanup()
|
|
|
|
def configure(self, payload):
|
|
rollcalc_app.CONVERSATION_SERVICE = ConversationService(
|
|
StaticNLUClient(payload),
|
|
store=InMemoryConversationStore(),
|
|
build_info=rollcalc_app.BUILD_INFO,
|
|
)
|
|
|
|
def create_conversation(self):
|
|
response = self.client.post("/api/conversations", headers=self.headers)
|
|
self.assertEqual(response.status_code, 201)
|
|
return response.get_json()["conversation_id"]
|
|
|
|
def test_conversation_success_exposes_authoritative_pdf(self):
|
|
self.configure({
|
|
"intent": "new_calculation",
|
|
"article_number": "180205",
|
|
"article_name_hint": "Bentofix NSP 4900",
|
|
"roll_length_m": 65.0,
|
|
"width_m": None,
|
|
"core_type": "150mm PVC",
|
|
"core_diameter_mm": None,
|
|
"include_roll_weight": False,
|
|
})
|
|
conversation_id = self.create_conversation()
|
|
response = self.client.post(
|
|
f"/api/conversations/{conversation_id}/messages",
|
|
json={"message": "Berechnen"},
|
|
headers=self.headers,
|
|
)
|
|
payload = response.get_json()
|
|
|
|
self.assertEqual(response.status_code, 200)
|
|
self.assertEqual(payload["status"], "success")
|
|
self.assertEqual(
|
|
payload["result"]["provenance"]["calculator"],
|
|
"roll_calculation.calculate_roll",
|
|
)
|
|
|
|
pdf = self.client.get(payload["pdf"]["url"], headers=self.headers)
|
|
self.assertEqual(pdf.status_code, 200)
|
|
self.assertEqual(pdf.mimetype, "application/pdf")
|
|
self.assertEqual(pdf.data.count(b"/Type /Page "), 1)
|
|
calculation = payload["result"]["calculation"]
|
|
self.assertIn(
|
|
f"{calculation['average_diameter_mm']:.1f} mm".encode("ascii"),
|
|
pdf.data,
|
|
)
|
|
|
|
def test_ollama_unavailable_returns_controlled_http_error(self):
|
|
self.configure(OllamaUnavailableError("offline"))
|
|
conversation_id = self.create_conversation()
|
|
|
|
response = self.client.post(
|
|
f"/api/conversations/{conversation_id}/messages",
|
|
json={"message": "Berechnen"},
|
|
headers=self.headers,
|
|
)
|
|
|
|
self.assertEqual(response.status_code, 503)
|
|
self.assertEqual(response.get_json()["status"], "nlu_error")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|