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()