route optimizer agent with dailygrubs ai assign
This commit is contained in:
0
tests/__init__.py
Normal file
0
tests/__init__.py
Normal file
142
tests/test_dispatch_agent.py
Normal file
142
tests/test_dispatch_agent.py
Normal file
@@ -0,0 +1,142 @@
|
||||
"""
|
||||
Unit tests for the DispatchAgent assignment-failure context gatherer.
|
||||
|
||||
Focus is the same as the ExceptionAgent tests: the gatherer must degrade
|
||||
gracefully (no coords, a failing GEORADIUS, a failing counter) and never raise.
|
||||
The real `_find_zone` is exercised through an in-memory fake Redis; the agent is
|
||||
built with __new__ so NATS / the message bus are never touched.
|
||||
|
||||
Run:
|
||||
python -m unittest discover -s tests
|
||||
"""
|
||||
import unittest
|
||||
|
||||
import nats.js.errors
|
||||
|
||||
from agents.dispatch_agent import DispatchAgent
|
||||
|
||||
|
||||
class FakeRedis:
|
||||
"""decode_responses=True style. `geo` is the GEORADIUS result for every call."""
|
||||
def __init__(self, geo=None, counter_start=0, raise_on=()):
|
||||
self._geo = list(geo) if geo is not None else []
|
||||
self._counter = counter_start
|
||||
self._raise_on = set(raise_on)
|
||||
self.incr_calls = []
|
||||
self.expire_calls = []
|
||||
|
||||
async def georadius(self, key, lon, lat, radius, unit, sort=None, count=None):
|
||||
if "georadius" in self._raise_on:
|
||||
raise RuntimeError("geo down")
|
||||
return list(self._geo)
|
||||
|
||||
async def hgetall(self, key):
|
||||
return {"hub_id": "H1", "zone_id": "z1", "avg_delivery_time": "60"}
|
||||
|
||||
async def incr(self, key):
|
||||
if "incr" in self._raise_on:
|
||||
raise RuntimeError("redis down")
|
||||
self._counter += 1
|
||||
self.incr_calls.append(key)
|
||||
return self._counter
|
||||
|
||||
async def expire(self, key, ttl):
|
||||
self.expire_calls.append((key, ttl))
|
||||
return True
|
||||
|
||||
|
||||
def make_agent(redis):
|
||||
agent = DispatchAgent.__new__(DispatchAgent) # skip heavy __init__ / message-bus registration
|
||||
agent._redis = redis
|
||||
return agent
|
||||
|
||||
|
||||
class FakeJS:
|
||||
"""Stand-in JetStream context for binding tests."""
|
||||
def __init__(self, stream_name=None, not_found=False):
|
||||
self._stream = stream_name
|
||||
self._not_found = not_found
|
||||
self.subscribe_calls = []
|
||||
self.add_stream_calls = []
|
||||
|
||||
async def find_stream_name_by_subject(self, subject):
|
||||
if self._not_found:
|
||||
raise nats.js.errors.NotFoundError()
|
||||
return self._stream
|
||||
|
||||
async def subscribe(self, subject, durable=None, stream=None, cb=None):
|
||||
self.subscribe_calls.append((subject, durable, stream))
|
||||
return object() # stand-in subscription
|
||||
|
||||
async def add_stream(self, **kwargs):
|
||||
self.add_stream_calls.append(kwargs)
|
||||
|
||||
|
||||
async def _cb(msg):
|
||||
pass
|
||||
|
||||
|
||||
def make_agent_js(js):
|
||||
agent = DispatchAgent.__new__(DispatchAgent)
|
||||
agent._nats_js = js
|
||||
return agent
|
||||
|
||||
|
||||
class TestGatherAssignmentFacts(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_rider_nearby(self):
|
||||
redis = FakeRedis(geo=["m5"], counter_start=0) # found on the first (10km) sweep
|
||||
facts, count = await make_agent(redis)._gather_assignment_facts("hyderabad", 17.4, 78.4)
|
||||
self.assertEqual(facts["zone_id"], "hyderabad")
|
||||
self.assertTrue(facts["has_coordinates"])
|
||||
self.assertEqual(facts["nearest_miler_within_km"], 10)
|
||||
self.assertEqual(facts["failures_today"], 1)
|
||||
self.assertEqual(count, 1)
|
||||
self.assertEqual(redis.expire_calls[0][1], 172800)
|
||||
|
||||
async def test_no_rider_within_30km(self):
|
||||
redis = FakeRedis(geo=[], counter_start=2) # every sweep empty
|
||||
facts, count = await make_agent(redis)._gather_assignment_facts("pune", 18.5, 73.8)
|
||||
self.assertEqual(facts["nearest_miler_within_km"], "none within 30km")
|
||||
self.assertEqual(count, 3)
|
||||
|
||||
async def test_no_coordinates(self):
|
||||
redis = FakeRedis(geo=["m1"])
|
||||
facts, count = await make_agent(redis)._gather_assignment_facts("unknown", None, None)
|
||||
self.assertFalse(facts["has_coordinates"])
|
||||
self.assertEqual(facts["nearest_miler_within_km"], "unknown (no coordinates)")
|
||||
self.assertEqual(count, 1) # counter still incremented
|
||||
self.assertEqual(redis.incr_calls and 1, 1)
|
||||
|
||||
async def test_georadius_error_is_swallowed(self):
|
||||
redis = FakeRedis(raise_on={"georadius"})
|
||||
facts, count = await make_agent(redis)._gather_assignment_facts("z", 1.0, 2.0) # must not raise
|
||||
self.assertEqual(facts["nearest_miler_within_km"], "none within 30km")
|
||||
self.assertEqual(facts["failures_today"], 1)
|
||||
|
||||
async def test_counter_error_is_swallowed(self):
|
||||
redis = FakeRedis(geo=["m9"], raise_on={"incr"})
|
||||
facts, count = await make_agent(redis)._gather_assignment_facts("z", 1.0, 2.0) # must not raise
|
||||
self.assertEqual(facts["nearest_miler_within_km"], 10) # coverage still gathered
|
||||
self.assertEqual(facts["failures_today"], 0) # counter degraded to 0
|
||||
self.assertEqual(count, 0)
|
||||
|
||||
|
||||
class TestBindConsumer(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_binds_to_discovered_stream_without_creating(self):
|
||||
js = FakeJS(stream_name="TRACKING")
|
||||
sub = await make_agent_js(js)._bind_consumer("booking.assigned", "d1", _cb)
|
||||
self.assertIsNotNone(sub)
|
||||
# Bound explicitly on the discovered stream...
|
||||
self.assertEqual(js.subscribe_calls, [("booking.assigned", "d1", "TRACKING")])
|
||||
# ...and never tried to create a stream (the old overlap bug).
|
||||
self.assertEqual(js.add_stream_calls, [])
|
||||
|
||||
async def test_returns_none_when_no_stream_carries_subject(self):
|
||||
js = FakeJS(not_found=True)
|
||||
sub = await make_agent_js(js)._bind_consumer("booking.assigned", "d1", _cb)
|
||||
self.assertIsNone(sub) # clear failure, not a raise
|
||||
self.assertEqual(js.subscribe_calls, []) # did not attempt to subscribe
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
74
tests/test_domain_agents.py
Normal file
74
tests/test_domain_agents.py
Normal file
@@ -0,0 +1,74 @@
|
||||
"""
|
||||
Tests for domain-agent correctness fixes.
|
||||
|
||||
FleetAgent: hub availability accounting must stay within [0, total] across
|
||||
maintenance / release / double-release (it used to drift because maintenance
|
||||
didn't decrement and release incremented unconditionally).
|
||||
|
||||
OrderAgent: validation/categorization must not crash on a null address block or
|
||||
a numeric pincode (valid JSON the backend may send).
|
||||
|
||||
Run:
|
||||
python -m unittest discover -s tests
|
||||
"""
|
||||
import unittest
|
||||
|
||||
from core.types import AgentTask
|
||||
from core.message_bus import message_bus
|
||||
from agents.fleet_agent import FleetAgent
|
||||
from agents.order_agent import OrderAgent
|
||||
|
||||
|
||||
def _task(task_type, **data):
|
||||
return AgentTask(task_id="t", agent_type="x", task_type=task_type, data=data)
|
||||
|
||||
|
||||
class TestFleetCapacity(unittest.IsolatedAsyncioTestCase):
|
||||
async def asyncTearDown(self):
|
||||
message_bus.unregister_agent("FLEET_AGENT")
|
||||
|
||||
async def test_maintenance_release_stay_in_bounds(self):
|
||||
fleet = FleetAgent()
|
||||
hub = "DL-HUB-01"
|
||||
total = fleet._hub_capacity[hub]["total"]
|
||||
start = fleet._hub_capacity[hub]["available"]
|
||||
|
||||
await fleet._schedule_maintenance(_task("schedule_maintenance", vehicle_id="DL-V-001"))
|
||||
self.assertEqual(fleet._hub_capacity[hub]["available"], start - 1) # decremented
|
||||
|
||||
await fleet._release_vehicle(_task("release_vehicle", vehicle_id="DL-V-001"))
|
||||
self.assertEqual(fleet._hub_capacity[hub]["available"], start) # restored
|
||||
|
||||
# Releasing an already-available vehicle must not push above total.
|
||||
await fleet._release_vehicle(_task("release_vehicle", vehicle_id="DL-V-001"))
|
||||
self.assertEqual(fleet._hub_capacity[hub]["available"], start)
|
||||
self.assertLessEqual(fleet._hub_capacity[hub]["available"], total)
|
||||
|
||||
|
||||
class TestOrderRobustness(unittest.IsolatedAsyncioTestCase):
|
||||
async def asyncTearDown(self):
|
||||
message_bus.unregister_agent("ORDER_AGENT")
|
||||
|
||||
def test_addr_pincode_tolerates_null_and_numeric(self):
|
||||
self.assertEqual(OrderAgent._addr_pincode({"pickup_address": None}, "pickup_address"), "")
|
||||
self.assertEqual(OrderAgent._addr_pincode({"pickup_address": {"pincode": 400001}}, "pickup_address"), "400001")
|
||||
self.assertEqual(OrderAgent._addr_pincode({}, "pickup_address"), "")
|
||||
|
||||
def test_valid_pincode_handles_int_and_none(self):
|
||||
agent = OrderAgent()
|
||||
self.assertTrue(agent._is_valid_pincode(400001)) # numeric, 6 digits
|
||||
self.assertFalse(agent._is_valid_pincode(None))
|
||||
self.assertFalse(agent._is_valid_pincode("12"))
|
||||
|
||||
def test_categorize_does_not_crash_on_bad_addresses(self):
|
||||
agent = OrderAgent()
|
||||
out = agent._categorize_order_data({
|
||||
"pickup_address": None, # explicit null
|
||||
"delivery_address": {"pincode": 400001}, # numeric pincode
|
||||
"items": [],
|
||||
})
|
||||
self.assertIn("zone_type", out)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
199
tests/test_exception_agent.py
Normal file
199
tests/test_exception_agent.py
Normal file
@@ -0,0 +1,199 @@
|
||||
"""
|
||||
Unit tests for the ExceptionAgent read-only context gatherer and its stall
|
||||
counter, plus the tz helpers.
|
||||
|
||||
Focus is the *degradation* behaviour: a missing pool, a failing query, or a
|
||||
Redis hiccup must never raise — the gatherer returns whatever facts it could
|
||||
collect so the decision still runs (skewing safe). No real Postgres/Redis/NATS
|
||||
is touched: the agent is built with __new__ and handed in-memory fakes.
|
||||
|
||||
Run:
|
||||
python -m unittest discover -s tests
|
||||
"""
|
||||
import unittest
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
from agents.exception_agent import ExceptionAgent, _parse_ts, _utcnow
|
||||
|
||||
|
||||
# ── Fakes ──────────────────────────────────────────────────────────────────
|
||||
|
||||
class FakePGConn:
|
||||
def __init__(self, row=None, raise_err=False):
|
||||
self._row = row
|
||||
self._raise = raise_err
|
||||
|
||||
async def fetchrow(self, query, *args):
|
||||
if self._raise:
|
||||
raise RuntimeError("pg down")
|
||||
return self._row
|
||||
|
||||
|
||||
class _AcquireCtx:
|
||||
def __init__(self, conn):
|
||||
self._conn = conn
|
||||
|
||||
async def __aenter__(self):
|
||||
return self._conn
|
||||
|
||||
async def __aexit__(self, *exc):
|
||||
return False
|
||||
|
||||
|
||||
class FakePGPool:
|
||||
def __init__(self, row=None, raise_err=False):
|
||||
self._conn = FakePGConn(row, raise_err)
|
||||
|
||||
def acquire(self):
|
||||
return _AcquireCtx(self._conn)
|
||||
|
||||
|
||||
class FakeRedis:
|
||||
"""decode_responses=True style — returns str values."""
|
||||
def __init__(self, movement=None, counter=None, raise_on=()):
|
||||
self._movement = movement or {}
|
||||
self._counter = counter # str or None
|
||||
self._raise_on = set(raise_on)
|
||||
self.incr_calls = []
|
||||
self.expire_calls = []
|
||||
|
||||
async def hgetall(self, key):
|
||||
if "hgetall" in self._raise_on:
|
||||
raise RuntimeError("redis down")
|
||||
return dict(self._movement)
|
||||
|
||||
async def get(self, key):
|
||||
if "get" in self._raise_on:
|
||||
raise RuntimeError("redis down")
|
||||
return self._counter
|
||||
|
||||
async def incr(self, key):
|
||||
if "incr" in self._raise_on:
|
||||
raise RuntimeError("redis down")
|
||||
self.incr_calls.append(key)
|
||||
self._counter = str(int(self._counter or 0) + 1)
|
||||
return int(self._counter)
|
||||
|
||||
async def expire(self, key, ttl):
|
||||
self.expire_calls.append((key, ttl))
|
||||
return True
|
||||
|
||||
|
||||
def make_agent(pg, redis):
|
||||
agent = ExceptionAgent.__new__(ExceptionAgent) # skip heavy __init__ / message-bus registration
|
||||
agent._pg = pg
|
||||
agent._redis = redis
|
||||
return agent
|
||||
|
||||
|
||||
# ── Tz helpers ───────────────────────────────────────────────────────────────
|
||||
|
||||
class TestTimeHelpers(unittest.TestCase):
|
||||
def test_utcnow_is_aware(self):
|
||||
self.assertIsNotNone(_utcnow().tzinfo)
|
||||
|
||||
def test_parse_ts_none_and_blank(self):
|
||||
self.assertIsNone(_parse_ts(None))
|
||||
self.assertIsNone(_parse_ts(""))
|
||||
|
||||
def test_parse_ts_invalid(self):
|
||||
self.assertIsNone(_parse_ts("not-a-timestamp"))
|
||||
|
||||
def test_parse_ts_naive_assumed_utc(self):
|
||||
dt = _parse_ts("2024-01-01T00:00:00")
|
||||
self.assertIsNotNone(dt)
|
||||
self.assertEqual(dt.utcoffset(), timedelta(0))
|
||||
|
||||
def test_parse_ts_aware_preserved(self):
|
||||
dt = _parse_ts("2024-01-01T00:00:00+00:00")
|
||||
self.assertEqual(dt.utcoffset(), timedelta(0))
|
||||
|
||||
|
||||
# ── Gatherer ───────────────────────────────────────────────────────────────
|
||||
|
||||
class TestGatherStallFacts(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_full_context(self):
|
||||
now = _utcnow()
|
||||
pg = FakePGPool(row={"status": "Miler_Assigned", "createdat": now - timedelta(minutes=30)})
|
||||
redis = FakeRedis(
|
||||
movement={
|
||||
"updated_at": (now - timedelta(minutes=5)).isoformat(),
|
||||
"position_unchanged_since": (now - timedelta(minutes=15)).isoformat(),
|
||||
},
|
||||
counter="3",
|
||||
)
|
||||
facts = await make_agent(pg, redis)._gather_stall_facts("m1", "b1")
|
||||
|
||||
self.assertEqual(facts["booking_status"], "Miler_Assigned")
|
||||
self.assertAlmostEqual(facts["minutes_since_booking_created"], 30.0, delta=0.2)
|
||||
self.assertAlmostEqual(facts["last_gps_ping_minutes_ago"], 5.0, delta=0.2)
|
||||
self.assertAlmostEqual(facts["position_unchanged_minutes"], 15.0, delta=0.2)
|
||||
self.assertEqual(facts["stalls_today"], 3)
|
||||
|
||||
async def test_naive_createdat_treated_as_utc(self):
|
||||
naive_created = _utcnow().replace(tzinfo=None) - timedelta(minutes=20) # naive UTC, as a `timestamp` column returns
|
||||
pg = FakePGPool(row={"status": "Pickup_Scheduled", "createdat": naive_created})
|
||||
facts = await make_agent(pg, FakeRedis())._gather_stall_facts("m1", "b1")
|
||||
self.assertAlmostEqual(facts["minutes_since_booking_created"], 20.0, delta=0.2) # no tz error
|
||||
|
||||
async def test_no_pg_pool(self):
|
||||
redis = FakeRedis(movement={"updated_at": _utcnow().isoformat()}, counter="1")
|
||||
facts = await make_agent(None, redis)._gather_stall_facts("m1", "b1")
|
||||
self.assertNotIn("booking_status", facts) # no booking facts
|
||||
self.assertIn("last_gps_ping_minutes_ago", facts) # redis facts still gathered
|
||||
self.assertEqual(facts["stalls_today"], 1)
|
||||
|
||||
async def test_pg_error_is_swallowed(self):
|
||||
pg = FakePGPool(raise_err=True)
|
||||
redis = FakeRedis(counter="2")
|
||||
facts = await make_agent(pg, redis)._gather_stall_facts("m1", "b1") # must not raise
|
||||
self.assertNotIn("booking_status", facts)
|
||||
self.assertEqual(facts["stalls_today"], 2)
|
||||
|
||||
async def test_redis_hgetall_error_is_swallowed(self):
|
||||
pg = FakePGPool(row={"status": "Miler_Assigned", "createdat": _utcnow()})
|
||||
redis = FakeRedis(raise_on={"hgetall"}, counter="1")
|
||||
facts = await make_agent(pg, redis)._gather_stall_facts("m1", "b1") # must not raise
|
||||
self.assertEqual(facts["booking_status"], "Miler_Assigned") # pg facts still gathered
|
||||
self.assertNotIn("position_unchanged_minutes", facts) # movement facts skipped
|
||||
self.assertEqual(facts["stalls_today"], 1) # counter still read
|
||||
|
||||
async def test_counter_read_error_is_swallowed(self):
|
||||
redis = FakeRedis(movement={"updated_at": _utcnow().isoformat()}, raise_on={"get"})
|
||||
facts = await make_agent(None, redis)._gather_stall_facts("m1", "b1") # must not raise
|
||||
self.assertIn("last_gps_ping_minutes_ago", facts)
|
||||
self.assertNotIn("stalls_today", facts)
|
||||
|
||||
async def test_empty_everywhere(self):
|
||||
pg = FakePGPool(row=None)
|
||||
facts = await make_agent(pg, FakeRedis())._gather_stall_facts("m1", "b1")
|
||||
self.assertEqual(facts, {}) # nothing available → empty, no crash
|
||||
|
||||
|
||||
# ── Stall counter ────────────────────────────────────────────────────────────
|
||||
|
||||
class TestStallCounter(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_counter_key_is_daily(self):
|
||||
agent = make_agent(None, FakeRedis())
|
||||
key = agent._stall_counter_key("m42")
|
||||
self.assertTrue(key.startswith("miler:m42:stalls:"))
|
||||
self.assertEqual(key.split(":")[-1], _utcnow().strftime("%Y-%m-%d"))
|
||||
|
||||
async def test_incr_sets_ttl(self):
|
||||
redis = FakeRedis(counter="0")
|
||||
agent = make_agent(None, redis)
|
||||
await agent._incr_stall_counter("m1")
|
||||
self.assertEqual(len(redis.incr_calls), 1)
|
||||
self.assertEqual(len(redis.expire_calls), 1)
|
||||
self.assertEqual(redis.expire_calls[0][1], 172800) # 48h TTL
|
||||
self.assertEqual(redis._counter, "1")
|
||||
|
||||
async def test_incr_error_is_swallowed(self):
|
||||
redis = FakeRedis(raise_on={"incr"})
|
||||
agent = make_agent(None, redis)
|
||||
await agent._incr_stall_counter("m1") # must not raise
|
||||
self.assertEqual(redis.incr_calls, [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
99
tests/test_http_client.py
Normal file
99
tests/test_http_client.py
Normal file
@@ -0,0 +1,99 @@
|
||||
"""
|
||||
Tests for core.http_client status handling.
|
||||
|
||||
Regression guard: the client used to return a truthy body for ANY status, so a
|
||||
4xx/5xx read as success. Now only 2xx/3xx yield a result; 4xx returns None with
|
||||
no retry; 5xx retries then returns None.
|
||||
|
||||
Run:
|
||||
python -m unittest discover -s tests
|
||||
"""
|
||||
import unittest
|
||||
from unittest.mock import patch, AsyncMock
|
||||
|
||||
import core.http_client as hc
|
||||
|
||||
|
||||
class FakeResp:
|
||||
def __init__(self, status, json_data=None, content_type="application/json", text=""):
|
||||
self.status = status
|
||||
self._json = json_data
|
||||
self.content_type = content_type
|
||||
self._text = text
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *a):
|
||||
return False
|
||||
|
||||
async def json(self):
|
||||
return self._json
|
||||
|
||||
async def text(self):
|
||||
return self._text
|
||||
|
||||
|
||||
class FakeSession:
|
||||
def __init__(self, responses):
|
||||
self._responses = list(responses)
|
||||
self.closed = False
|
||||
self.calls = []
|
||||
|
||||
def _make(self, method, url, **kwargs):
|
||||
self.calls.append((method, url))
|
||||
return self._responses.pop(0)
|
||||
|
||||
def get(self, url, **kwargs):
|
||||
return self._make("get", url, **kwargs)
|
||||
|
||||
def post(self, url, **kwargs):
|
||||
return self._make("post", url, **kwargs)
|
||||
|
||||
def patch(self, url, **kwargs):
|
||||
return self._make("patch", url, **kwargs)
|
||||
|
||||
|
||||
class TestHttpStatus(unittest.IsolatedAsyncioTestCase):
|
||||
def _install(self, responses):
|
||||
session = FakeSession(responses)
|
||||
hc._session = session
|
||||
return session
|
||||
|
||||
async def asyncTearDown(self):
|
||||
hc._session = None
|
||||
|
||||
async def test_2xx_json_returned(self):
|
||||
self._install([FakeResp(200, json_data={"ok": True})])
|
||||
result = await hc.api_post("http://x/act")
|
||||
self.assertEqual(result, {"ok": True})
|
||||
|
||||
async def test_4xx_returns_none_no_retry(self):
|
||||
session = self._install([FakeResp(404, text="not found")])
|
||||
result = await hc.api_post("http://x/missing")
|
||||
self.assertIsNone(result) # rejection is not success
|
||||
self.assertEqual(len(session.calls), 1) # 4xx is not retried
|
||||
|
||||
async def test_5xx_retries_then_none(self):
|
||||
session = self._install([FakeResp(500, text="boom")] * 3)
|
||||
with patch.object(hc.asyncio, "sleep", new=AsyncMock()):
|
||||
result = await hc.api_post("http://x/act")
|
||||
self.assertIsNone(result)
|
||||
self.assertEqual(len(session.calls), 3) # retried up to max
|
||||
|
||||
async def test_5xx_then_2xx_recovers(self):
|
||||
session = self._install([FakeResp(503, text="try later"),
|
||||
FakeResp(200, json_data={"ok": 1})])
|
||||
with patch.object(hc.asyncio, "sleep", new=AsyncMock()):
|
||||
result = await hc.api_post("http://x/act")
|
||||
self.assertEqual(result, {"ok": 1})
|
||||
self.assertEqual(len(session.calls), 2)
|
||||
|
||||
async def test_2xx_non_json_returns_status(self):
|
||||
self._install([FakeResp(204, content_type="text/plain")])
|
||||
result = await hc.api_post("http://x/act")
|
||||
self.assertEqual(result, {"status_code": 204})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
94
tests/test_llm.py
Normal file
94
tests/test_llm.py
Normal file
@@ -0,0 +1,94 @@
|
||||
"""
|
||||
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()
|
||||
75
tests/test_message_bus.py
Normal file
75
tests/test_message_bus.py
Normal file
@@ -0,0 +1,75 @@
|
||||
"""
|
||||
Tests for directed-message delivery on the MessageBus.
|
||||
|
||||
Regression guard for the dead-letter bug: directed messages used to be appended
|
||||
to an internal queue that nothing ever drained. They must now be delivered live
|
||||
to the recipient agent — task-bearing payloads become tasks, notifications reach
|
||||
handle_message, and an absent recipient still retains the message in the pull
|
||||
queue (never silently lost).
|
||||
|
||||
Run:
|
||||
python -m unittest discover -s tests
|
||||
"""
|
||||
import unittest
|
||||
|
||||
from core.agent import Agent
|
||||
from core.message_bus import message_bus
|
||||
from core.types import MessageType
|
||||
|
||||
|
||||
class Recorder(Agent):
|
||||
"""Minimal concrete agent that records what it was handed."""
|
||||
def __init__(self, agent_id):
|
||||
super().__init__(agent_id, "test", "recorder") # registers with the bus
|
||||
self.notifications = []
|
||||
|
||||
async def handle_task(self, task):
|
||||
return {"ok": True}
|
||||
|
||||
async def handle_message(self, message):
|
||||
self.notifications.append(message)
|
||||
|
||||
|
||||
class TestDirectedDelivery(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_task_message_enqueued_as_task(self):
|
||||
agent = Recorder("REC_TASK")
|
||||
try:
|
||||
await message_bus.send_to_agent(
|
||||
sender="X", recipient="REC_TASK",
|
||||
message_type=MessageType.AGENT_TASK,
|
||||
payload={"task_type": "do_thing", "booking_id": "b1"},
|
||||
)
|
||||
task = agent._task_queue.get_nowait() # delivered, not dead-lettered
|
||||
self.assertEqual(task.task_type, "do_thing")
|
||||
self.assertEqual(task.data["booking_id"], "b1")
|
||||
finally:
|
||||
message_bus.unregister_agent("REC_TASK")
|
||||
|
||||
async def test_notification_goes_to_handle_message(self):
|
||||
agent = Recorder("REC_NOTIF")
|
||||
try:
|
||||
await message_bus.send_to_agent(
|
||||
sender="X", recipient="REC_NOTIF",
|
||||
message_type=MessageType.EXCEPTION_DETECTED,
|
||||
payload={"order_id": "b9"}, # no task_type
|
||||
)
|
||||
self.assertEqual(len(agent.notifications), 1)
|
||||
self.assertEqual(agent.notifications[0].message_type, MessageType.EXCEPTION_DETECTED)
|
||||
self.assertTrue(agent._task_queue.empty()) # a notification is not a task
|
||||
finally:
|
||||
message_bus.unregister_agent("REC_NOTIF")
|
||||
|
||||
async def test_absent_recipient_retained_in_pull_queue(self):
|
||||
# Nobody registered under this id -> message kept, not lost.
|
||||
await message_bus.send_to_agent(
|
||||
sender="X", recipient="NOBODY_HOME",
|
||||
message_type=MessageType.AGENT_TASK,
|
||||
payload={"task_type": "x"},
|
||||
)
|
||||
msgs = await message_bus.get_messages("NOBODY_HOME")
|
||||
self.assertEqual(len(msgs), 1)
|
||||
self.assertEqual(msgs[0].recipient, "NOBODY_HOME")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
66
tests/test_stall_dedup.py
Normal file
66
tests/test_stall_dedup.py
Normal file
@@ -0,0 +1,66 @@
|
||||
"""
|
||||
Tests for the Redis-backed stall dedup claim (replaces the old in-memory set).
|
||||
|
||||
- The first claim for a booking wins; a second within the window is refused.
|
||||
- The consumer-side handled-claim behaves the same (redelivery guard).
|
||||
- A Redis error fails OPEN (returns True) so a genuine stall is never muted.
|
||||
|
||||
Run:
|
||||
python -m unittest discover -s tests
|
||||
"""
|
||||
import unittest
|
||||
|
||||
from agents.exception_agent import ExceptionAgent
|
||||
|
||||
|
||||
class FakeRedisNX:
|
||||
"""Minimal Redis with SET NX + TTL semantics."""
|
||||
def __init__(self, raise_on_set=False):
|
||||
self.store = {}
|
||||
self.raise_on_set = raise_on_set
|
||||
self.set_calls = []
|
||||
|
||||
async def set(self, key, val, ex=None, nx=False):
|
||||
if self.raise_on_set:
|
||||
raise RuntimeError("redis down")
|
||||
self.set_calls.append((key, ex, nx))
|
||||
if nx and key in self.store:
|
||||
return None # already claimed
|
||||
self.store[key] = val
|
||||
return True
|
||||
|
||||
|
||||
def make_agent(redis):
|
||||
agent = ExceptionAgent.__new__(ExceptionAgent) # skip heavy __init__
|
||||
agent._redis = redis
|
||||
return agent
|
||||
|
||||
|
||||
class TestStallDedup(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_first_alert_claim_wins_second_refused(self):
|
||||
agent = make_agent(FakeRedisNX())
|
||||
self.assertTrue(await agent._claim_stall("b1")) # first wins
|
||||
self.assertFalse(await agent._claim_stall("b1")) # duplicate refused
|
||||
self.assertTrue(await agent._claim_stall("b2")) # different booking ok
|
||||
|
||||
async def test_claim_sets_ttl_and_nx(self):
|
||||
redis = FakeRedisNX()
|
||||
await make_agent(redis)._claim_stall("b9")
|
||||
key, ex, nx = redis.set_calls[0]
|
||||
self.assertEqual(key, "stall_notified:b9")
|
||||
self.assertTrue(nx)
|
||||
self.assertEqual(ex, 21600)
|
||||
|
||||
async def test_handled_claim_guards_redelivery(self):
|
||||
agent = make_agent(FakeRedisNX())
|
||||
self.assertTrue(await agent._claim_stall_handled("b1"))
|
||||
self.assertFalse(await agent._claim_stall_handled("b1"))
|
||||
|
||||
async def test_redis_error_fails_open(self):
|
||||
agent = make_agent(FakeRedisNX(raise_on_set=True))
|
||||
self.assertTrue(await agent._claim_stall("b1")) # must not mute a real stall
|
||||
self.assertTrue(await agent._claim_stall_handled("b1"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user