Files
AI_engine/tests/test_llm.py

95 lines
3.0 KiB
Python

"""
Tests for core.llm._decide plumbing with a mocked async client (no real API).
Covers: a valid structured decision is parsed; refusal/truncation/invalid-action
return None (-> deterministic fallback); a transient error retries once and can
recover; a persistent error gives up after the retry.
Run:
python -m unittest discover -s tests
"""
import json
import unittest
from unittest.mock import patch
import core.llm as llm
class Block:
def __init__(self, btype, text=None):
self.type = btype
self.text = text
class FakeResp:
def __init__(self, content, stop_reason="end_turn"):
self.content = content
self.stop_reason = stop_reason
class FakeMessages:
def __init__(self, resp=None, exc=None, exc_then=None):
self._resp = resp
self._exc = exc
self._exc_then = exc_then # raise on first call, then succeed
self.calls = 0
async def create(self, **kwargs):
self.calls += 1
if self._exc_then is not None and self.calls == 1:
raise self._exc_then
if self._exc is not None:
raise self._exc
return self._resp
class FakeClient:
def __init__(self, messages):
self.messages = messages
def _text_resp(action, stop="end_turn"):
payload = json.dumps({"action": action, "reasoning": "because", "confidence": 0.8})
return FakeResp([Block("thinking"), Block("text", payload)], stop_reason=stop)
class TestDecide(unittest.IsolatedAsyncioTestCase):
async def _run_stall(self, messages):
with patch.object(llm, "_get_client", return_value=FakeClient(messages)):
return await llm.decide_stall_response("ctx")
async def test_valid_decision_parsed(self):
d = await self._run_stall(FakeMessages(resp=_text_resp("wait")))
self.assertIsNotNone(d)
self.assertEqual(d.action, "wait")
self.assertEqual(d.confidence, 0.8)
async def test_refusal_returns_none(self):
d = await self._run_stall(FakeMessages(resp=_text_resp("wait", stop="refusal")))
self.assertIsNone(d)
async def test_truncation_returns_none(self):
d = await self._run_stall(FakeMessages(resp=_text_resp("wait", stop="max_tokens")))
self.assertIsNone(d)
async def test_invalid_action_returns_none(self):
d = await self._run_stall(FakeMessages(resp=_text_resp("teleport")))
self.assertIsNone(d)
async def test_transient_error_retries_then_succeeds(self):
msgs = FakeMessages(resp=_text_resp("escalate"), exc_then=RuntimeError("blip"))
d = await self._run_stall(msgs)
self.assertIsNotNone(d)
self.assertEqual(d.action, "escalate")
self.assertEqual(msgs.calls, 2) # retried once
async def test_persistent_error_gives_up(self):
msgs = FakeMessages(exc=RuntimeError("down"))
d = await self._run_stall(msgs)
self.assertIsNone(d)
self.assertEqual(msgs.calls, 2) # initial + one retry, then stop
if __name__ == "__main__":
unittest.main()