main
py 171 lines 4.98 KB
Raw
1 import asyncio
2 import sys
3 from pathlib import Path
4 from types import SimpleNamespace
5
6 import pytest
7
8 PROJECT_ROOT = Path(__file__).resolve().parents[1]
9 if str(PROJECT_ROOT) not in sys.path:
10 sys.path.insert(0, str(PROJECT_ROOT))
11
12 from agent import LoopData
13 from plugins._memory.extensions.python.message_loop_prompts_after import (
14 _50_recall_memories as recall_module,
15 _91_recall_wait as wait_module,
16 )
17
18
19 class _LogItem:
20 def update(self, **_kwargs):
21 pass
22
23
24 class _Agent:
25 def __init__(self):
26 self.data = {}
27 self.project = "project-a"
28 self.config = SimpleNamespace(profile="agent0")
29 self.context = SimpleNamespace(
30 log=SimpleNamespace(log=lambda **_kwargs: _LogItem()),
31 get_data=lambda *_args, **_kwargs: self.project,
32 )
33
34 def get_data(self, key):
35 return self.data.get(key)
36
37 def set_data(self, key, value):
38 self.data[key] = value
39
40 def read_prompt(self, _name):
41 return "Recall is running in the background."
42
43
44 def _settings():
45 return {
46 "memory_recall_enabled": True,
47 "memory_recall_delayed": True,
48 "memory_recall_interval": 1,
49 }
50
51
52 @pytest.mark.asyncio
53 async def test_delayed_recall_result_reaches_the_next_monologue(monkeypatch):
54 agent = _Agent()
55 recall = recall_module.RecallMemories(agent=agent)
56 monkeypatch.setattr(
57 recall_module.plugins, "get_plugin_config", lambda *_args: _settings()
58 )
59
60 async def search_memories(**_kwargs):
61 return {"memories": "recalled context"}
62
63 monkeypatch.setattr(recall, "search_memories", search_memories)
64
65 first_loop = LoopData()
66 first_loop.iteration = 0
67 await recall.execute(loop_data=first_loop)
68 await agent.get_data(recall_module.DATA_NAME_TASK)
69
70 next_loop = LoopData()
71 next_loop.iteration = 0
72 await recall.execute(loop_data=next_loop)
73 await agent.get_data(recall_module.DATA_NAME_TASK)
74
75 assert next_loop.extras_persistent["memories"] == "recalled context"
76
77
78 @pytest.mark.asyncio
79 async def test_delayed_recall_task_survives_the_next_internal_iteration(monkeypatch):
80 agent = _Agent()
81 recall = recall_module.RecallMemories(agent=agent)
82 wait = wait_module.RecallWait(agent=agent)
83 settings = _settings()
84 monkeypatch.setattr(
85 recall_module.plugins, "get_plugin_config", lambda *_args: settings
86 )
87
88 release = asyncio.Event()
89
90 async def search_memories(**_kwargs):
91 await release.wait()
92 return {"solutions": "recalled solution"}
93
94 monkeypatch.setattr(recall, "search_memories", search_memories)
95
96 loop_data = LoopData()
97 loop_data.iteration = 0
98 await recall.execute(loop_data=loop_data)
99 first_task = agent.get_data(recall_module.DATA_NAME_TASK)
100 await wait.execute(loop_data=loop_data)
101 assert "memory_recall_delayed" in loop_data.extras_temporary
102
103 loop_data.iteration = 1
104 await recall.execute(loop_data=loop_data)
105 next_task = agent.get_data(recall_module.DATA_NAME_TASK)
106 release.set()
107 await first_task
108 if next_task is not first_task:
109 await next_task
110
111 assert next_task is first_task
112 await wait.execute(loop_data=loop_data)
113 assert loop_data.extras_persistent["solutions"] == "recalled solution"
114
115
116 @pytest.mark.asyncio
117 async def test_completed_blocking_recall_is_applied(monkeypatch):
118 agent = _Agent()
119 recall = recall_module.RecallMemories(agent=agent)
120 wait = wait_module.RecallWait(agent=agent)
121 settings = {**_settings(), "memory_recall_delayed": False}
122 monkeypatch.setattr(
123 recall_module.plugins, "get_plugin_config", lambda *_args: settings
124 )
125
126 async def search_memories(**_kwargs):
127 return {"memories": "ready before wait"}
128
129 monkeypatch.setattr(recall, "search_memories", search_memories)
130
131 loop_data = LoopData()
132 loop_data.iteration = 0
133 await recall.execute(loop_data=loop_data)
134 await agent.get_data(recall_module.DATA_NAME_TASK)
135 await wait.execute(loop_data=loop_data)
136
137 assert loop_data.extras_persistent["memories"] == "ready before wait"
138
139
140 @pytest.mark.asyncio
141 async def test_delayed_recall_result_does_not_cross_profile_or_project(monkeypatch):
142 agent = _Agent()
143 recall = recall_module.RecallMemories(agent=agent)
144 monkeypatch.setattr(
145 recall_module.plugins, "get_plugin_config", lambda *_args: _settings()
146 )
147
148 release = asyncio.Event()
149
150 async def search_memories(**_kwargs):
151 await release.wait()
152 return {"memories": "project-a memory"}
153
154 monkeypatch.setattr(recall, "search_memories", search_memories)
155
156 first_loop = LoopData()
157 first_loop.iteration = 0
158 await recall.execute(loop_data=first_loop)
159 first_task = agent.get_data(recall_module.DATA_NAME_TASK)
160
161 agent.project = "project-b"
162 agent.config.profile = "developer"
163 release.set()
164 await first_task
165
166 next_loop = LoopData()
167 next_loop.iteration = 0
168 await recall.execute(loop_data=next_loop)
169
170 assert "memories" not in next_loop.extras_persistent
171 await agent.get_data(recall_module.DATA_NAME_TASK)