import unittest from unittest.mock import patch import api.state as state _TERMINAL_STATES = ( "SUCCESS", "COMPLETED", "ERROR", "CANCELLED", "RATE_LIMITED", ) class AgentTaskPruningTests(unittest.TestCase): def test_prune_expires_every_terminal_state_and_cleans_byok(self): now_ms = 10_000_000 old_ms = now_ms - state._AGENT_TASK_TTL_MS - 1 task_ids = {f"old-{status.lower()}" for status in _TERMINAL_STATES} task_ids.update({"running", "queued"}) tasks = { task_id: { "status": ( task_id.removeprefix("old-").upper() if task_id.startswith("old-") else task_id.upper() ), "created_at": old_ms, } for task_id in task_ids } byok_clients = {task_id: object() for task_id in task_ids} with patch.object(state, "_agent_tasks", tasks), patch.object( state, "_task_ai_clients", byok_clients ), patch.object(state.time, "time", return_value=now_ms / 1000): state._prune_agent_tasks() self.assertEqual(set(state._agent_tasks), {"running", "queued"}) self.assertEqual(set(state._task_ai_clients), {"running", "queued"}) def test_prune_keeps_recent_terminal_tasks(self): now_ms = 10_000_000 recent_ms = now_ms - state._AGENT_TASK_TTL_MS + 1 tasks = { status.lower(): {"status": status, "created_at": recent_ms} for status in _TERMINAL_STATES } with patch.object(state, "_agent_tasks", tasks), patch.object( state, "_task_ai_clients", {} ), patch.object(state.time, "time", return_value=now_ms / 1000): state._prune_agent_tasks() self.assertEqual( set(state._agent_tasks), {status.lower() for status in _TERMINAL_STATES}, ) def test_terminal_state_set_matches_pruning_contract(self): self.assertEqual( state._AGENT_TASK_TERMINAL_STATES, frozenset(_TERMINAL_STATES), ) self.assertNotIn("QUEUED", state._AGENT_TASK_TERMINAL_STATES) self.assertNotIn("RUNNING", state._AGENT_TASK_TERMINAL_STATES) self.assertNotIn("CREATING", state._AGENT_TASK_TERMINAL_STATES) if __name__ == "__main__": unittest.main()