39 lines
1.1 KiB
Python
39 lines
1.1 KiB
Python
import unittest
|
|
|
|
from nlprog.json_repair import complete_json, parse_json_object
|
|
from nlprog.llm import Message
|
|
|
|
|
|
class BrokenThenGoodClient:
|
|
def __init__(self):
|
|
self.calls = 0
|
|
|
|
def complete(self, messages):
|
|
self.calls += 1
|
|
if self.calls == 1:
|
|
return "not json"
|
|
return '{"final": "repaired"}'
|
|
|
|
|
|
class JsonRepairTests(unittest.TestCase):
|
|
def test_parse_json_object_strips_fences(self):
|
|
data, error = parse_json_object('```json\n{"ok": true}\n```')
|
|
self.assertEqual(error, "")
|
|
self.assertEqual(data, {"ok": True})
|
|
|
|
def test_complete_json_repairs_bad_response(self):
|
|
result = complete_json(
|
|
BrokenThenGoodClient(),
|
|
[Message("system", "test"), Message("user", "return json")],
|
|
'{"final": "summary"}',
|
|
lambda data: None if "final" in data else "missing final",
|
|
max_retries=2,
|
|
)
|
|
self.assertEqual(result.data, {"final": "repaired"})
|
|
self.assertEqual(result.attempts, 2)
|
|
self.assertTrue(result.repaired)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|