Spaces:
Running
Running
| 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() | |