Spaces:
Running
Running
File size: 2,415 Bytes
9f9d3dc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 | 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()
|