95 lines
3.0 KiB
Python
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()
|