Initial import of NLProg
This commit is contained in:
@@ -0,0 +1,38 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user