feat: add durable chat storage and disconnect recovery
This commit is contained in:
@@ -14,6 +14,13 @@ class ValidationHarnessTests(unittest.TestCase):
|
||||
self.assertEqual(validators, ["validator-a", "validator-b"])
|
||||
self.assertTrue(harness)
|
||||
|
||||
def test_sqlite_is_default_even_when_postgres_is_installed(self):
|
||||
with patch.object(api, "_storage_config", return_value={"backend": "sqlite"}), patch.object(
|
||||
api, "_postgres_configured", return_value=True
|
||||
):
|
||||
self.assertEqual(api._storage_backend(), "sqlite")
|
||||
self.assertFalse(api._postgres_enabled())
|
||||
|
||||
def test_harness_requires_one_distinct_validator(self):
|
||||
body = api.ChatRequest(primary_model="primary", validator_models=[], harness=True)
|
||||
with patch.object(api, "_require_installed_model", side_effect=lambda name: name):
|
||||
@@ -29,45 +36,28 @@ class ValidationHarnessTests(unittest.TestCase):
|
||||
self.assertEqual(validators, ["validator-a"])
|
||||
self.assertTrue(harness)
|
||||
|
||||
def test_chat_route_returns_one_compiled_answer(self):
|
||||
body = api.ChatRequest(
|
||||
primary_model="primary",
|
||||
validator_models=["validator-a", "validator-b"],
|
||||
harness=True,
|
||||
message="Answer this once and validate it.",
|
||||
)
|
||||
persisted_messages = []
|
||||
|
||||
def fake_stream(payload, request_id, cancel_event=None, parent_id=None):
|
||||
prompt = payload["messages"][-1]["content"]
|
||||
if payload["model"] == "primary" and "VALIDATION REPORTS:" not in prompt:
|
||||
content = "draft"
|
||||
elif payload["model"].startswith("validator"):
|
||||
content = "No material issues found."
|
||||
else:
|
||||
content = "single compiled answer"
|
||||
return {"message": {"role": "assistant", "content": content}, "done": True}
|
||||
|
||||
def fake_metric(conversation_id, request_id, model, state, status=None, error=""):
|
||||
return {"request_id": request_id, "model": model, "status": status}
|
||||
|
||||
with patch.object(api, "_harness_models", return_value=("primary", ["validator-a", "validator-b"], True)), patch.object(
|
||||
api, "_require_installed_model", side_effect=lambda name: name
|
||||
), patch.object(api, "_ensure_conversation"), patch.object(api, "_persist_message", side_effect=lambda *args, **kwargs: persisted_messages.append(args)), patch.object(
|
||||
api, "_persist_metric", side_effect=fake_metric
|
||||
), patch.object(api, "_stream_chat_request", side_effect=fake_stream), patch.object(
|
||||
api, "_runtime_snapshot", return_value={}
|
||||
):
|
||||
def test_chat_route_queues_server_owned_job(self):
|
||||
body = api.ChatRequest(primary_model="primary", message="Queue this request")
|
||||
fake_job = {
|
||||
"request_id": "request",
|
||||
"conversation_id": "conversation",
|
||||
"status": "queued",
|
||||
"mode": "direct",
|
||||
"primary_model": "primary",
|
||||
"validator_models_json": "[]",
|
||||
"result_json": "{}",
|
||||
"error": "",
|
||||
"updated_at": 1.0,
|
||||
}
|
||||
with patch.object(api, "_harness_models", return_value=("primary", [], False)), patch.object(
|
||||
api, "_get_chat_job", side_effect=[None, fake_job]
|
||||
), patch.object(api, "_ensure_conversation"), patch.object(api, "_persist_message"), patch.object(
|
||||
api, "_chat_state"
|
||||
), patch.object(api, "_create_chat_job"), patch.object(api, "_submit_chat_job"):
|
||||
response = api.chat(body)
|
||||
|
||||
self.assertEqual(response["mode"], "harness")
|
||||
self.assertEqual(response["message"]["content"], "single compiled answer")
|
||||
self.assertEqual(response["primary_model"], "primary")
|
||||
self.assertEqual(response["validator_models"], ["validator-a", "validator-b"])
|
||||
assistant_messages = [args for args in persisted_messages if len(args) >= 3 and args[2] == "assistant"]
|
||||
self.assertEqual(len(assistant_messages), 1)
|
||||
self.assertEqual(assistant_messages[0][3], "single compiled answer")
|
||||
self.assertEqual(len(response["validation_reports"]), 2)
|
||||
self.assertEqual(response["status"], "queued")
|
||||
self.assertFalse(response["done"])
|
||||
self.assertEqual(response["mode"], "direct")
|
||||
|
||||
def test_primary_draft_validators_and_primary_compilation_produce_one_answer(self):
|
||||
body = api.ChatRequest(
|
||||
@@ -95,7 +85,9 @@ class ValidationHarnessTests(unittest.TestCase):
|
||||
|
||||
with patch.object(api, "_require_installed_model", side_effect=lambda name: name), patch.object(
|
||||
api, "_stream_chat_request", side_effect=fake_stream
|
||||
), patch.object(api, "_persist_metric", side_effect=fake_metric):
|
||||
), patch.object(api, "_persist_metric", side_effect=fake_metric), patch.object(
|
||||
api, "_persist_chat_stage"
|
||||
), patch.object(api, "_persist_chat_event"):
|
||||
final, metrics, reports = api._run_validation_harness(
|
||||
body,
|
||||
"root-request",
|
||||
@@ -112,10 +104,6 @@ class ValidationHarnessTests(unittest.TestCase):
|
||||
self.assertEqual(len(calls), 4)
|
||||
self.assertEqual(calls[0][0], "primary")
|
||||
self.assertTrue(all(call[2] == "root-request" for call in calls))
|
||||
with api._chat_requests_lock:
|
||||
child_states = [state for key, state in api._chat_requests.items() if key != "root-request"]
|
||||
self.assertGreaterEqual(len(child_states), 4)
|
||||
self.assertTrue(all(state.get("parent_id") == "root-request" for state in child_states[-4:]))
|
||||
compiler_prompt = calls[-1][1]
|
||||
self.assertIn("PRIMARY DRAFT:", compiler_prompt)
|
||||
self.assertIn("VALIDATOR 1", compiler_prompt)
|
||||
|
||||
Reference in New Issue
Block a user