File size: 11,071 Bytes
8f7a08e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
"""
Tests for memory system and temporal reasoning.
"""

import pytest
from app.env import CRMQueryEnv
from app.tasks import get_task_by_id
from app.models import Observation


class TestMemoryInitialization:
    """Test memory system initialization."""

    def test_memory_fields_in_state(self):
        """Test that environment state has memory fields."""
        env = CRMQueryEnv()
        obs = env.reset()
        
        assert hasattr(obs, 'memory_cache')
        assert hasattr(obs, 'step_summaries')
        assert isinstance(obs.memory_cache, dict)
        assert isinstance(obs.step_summaries, list)

    def test_retrieved_entities_initialization(self):
        """Test that retrieved entities are initialized."""
        env = CRMQueryEnv()
        env.reset()
        
        assert 'customers' in env.retrieved_entities
        assert 'orders' in env.retrieved_entities
        assert 'tickets' in env.retrieved_entities
        assert len(env.retrieved_entities['customers']) == 0
        assert len(env.retrieved_entities['orders']) == 0
        assert len(env.retrieved_entities['tickets']) == 0

    def test_step_summaries_initialization(self):
        """Test that step summaries are initialized."""
        env = CRMQueryEnv()
        env.reset()
        
        assert isinstance(env.step_summaries, list)
        assert len(env.step_summaries) == 0


class TestEntityCaching:
    """Test entity caching and retrieval."""

    def test_customers_cached_on_search(self):
        """Test that customers are cached when searched."""
        env = CRMQueryEnv()
        env.reset()
        
        action = {
            "tool": "search_customers",
            "arguments": {"tier": "Gold"}
        }
        
        obs, reward, done, info = env.step(action)
        
        assert len(env.retrieved_entities['customers']) > 0
        assert len(obs.memory_cache['customers']) > 0

    def test_orders_cached_on_search(self):
        """Test that orders are cached when searched."""
        env = CRMQueryEnv()
        env.reset()
        
        action = {
            "tool": "search_orders",
            "arguments": {"status": "Completed"}
        }
        
        obs, reward, done, info = env.step(action)
        
        assert len(env.retrieved_entities['orders']) > 0
        assert len(obs.memory_cache['orders']) > 0

    def test_tickets_cached_on_search(self):
        """Test that tickets are cached when searched."""
        env = CRMQueryEnv()
        env.reset()
        
        action = {
            "tool": "search_tickets",
            "arguments": {"priority": "High"}
        }
        
        obs, reward, done, info = env.step(action)
        
        assert len(env.retrieved_entities['tickets']) > 0
        assert len(obs.memory_cache['tickets']) > 0

    def test_multiple_queries_accumulate(self):
        """Test that multiple queries accumulate cached entities."""
        env = CRMQueryEnv()
        env.reset()
        
        # First query
        action1 = {"tool": "search_customers", "arguments": {"tier": "Gold"}}
        env.step(action1)
        count1 = len(env.retrieved_entities['customers'])
        
        # Second query
        action2 = {"tool": "search_customers", "arguments": {"tier": "Silver"}}
        env.step(action2)
        count2 = len(env.retrieved_entities['customers'])
        
        assert count2 >= count1
        assert count2 > 0


class TestStepSummaries:
    """Test step summary generation."""

    def test_summary_created_per_step(self):
        """Test that summary is created for each step."""
        env = CRMQueryEnv()
        env.reset()
        
        assert len(env.step_summaries) == 0
        
        action = {
            "tool": "search_customers",
            "arguments": {"tier": "Gold"}
        }
        env.step(action)
        
        assert len(env.step_summaries) == 1
        assert isinstance(env.step_summaries[0], str)
        assert "search_customers" in env.step_summaries[0]

    def test_summary_format(self):
        """Test that summaries have expected format."""
        env = CRMQueryEnv()
        env.reset()
        
        action = {
            "tool": "search_orders",
            "arguments": {"product": "Laptop"}
        }
        env.step(action)
        
        summary = env.step_summaries[0]
        assert "Step" in summary
        assert "search_orders" in summary
        assert "results" in summary

    def test_multiple_summaries_preserved(self):
        """Test that multiple step summaries are preserved."""
        env = CRMQueryEnv()
        env.reset()
        
        actions = [
            {"tool": "search_customers", "arguments": {"tier": "Gold"}},
            {"tool": "search_tickets", "arguments": {"priority": "High"}},
            {"tool": "search_orders", "arguments": {"status": "Completed"}},
        ]
        
        for action in actions:
            env.step(action)
        
        assert len(env.step_summaries) == 3
        assert "search_customers" in env.step_summaries[0]
        assert "search_tickets" in env.step_summaries[1]
        assert "search_orders" in env.step_summaries[2]


class TestMemoryReuseRewards:
    """Test memory reuse reward components."""

    def test_memory_reuse_bonus(self):
        """Test that memory reuse is tracked and rewarded."""
        env = CRMQueryEnv()
        env.reset()
        
        # First query
        action1 = {
            "tool": "search_customers",
            "arguments": {"tier": "Gold"}
        }
        obs1, reward1, done1, info1 = env.step(action1)
        
        # Second query (different filter)
        action2 = {
            "tool": "search_customers",
            "arguments": {"tier": "Silver"}
        }
        obs2, reward2, done2, info2 = env.step(action2)
        
        # Both should have valid rewards
        assert reward1.value > -10
        assert reward2.value > -10

    def test_cache_maintained_component(self):
        """Test cache maintained reward component."""
        env = CRMQueryEnv()
        env.reset()
        
        # Multiple searches build cache
        action1 = {"tool": "search_customers", "arguments": {"tier": "Gold"}}
        obs1, reward1, done1, info1 = env.step(action1)
        
        action2 = {"tool": "search_tickets", "arguments": {"priority": "High"}}
        obs2, reward2, done2, info2 = env.step(action2)
        
        # Verify cache is growing
        assert len(env.retrieved_entities['customers']) > 0
        assert len(env.retrieved_entities['tickets']) > 0

    def test_memory_hit_tracking(self):
        """Test that memory hits are tracked in history."""
        env = CRMQueryEnv()
        env.reset()
        
        action = {"tool": "search_customers", "arguments": {"tier": "Gold"}}
        env.step(action)
        
        # Check history tracks memory hit
        assert len(env.history) == 1
        assert "memory_hit" in env.history[0]


class TestRedundantQueryPenalties:
    """Test penalties for redundant queries."""

    def test_repeated_query_penalty(self):
        """Test that repeated queries receive penalties."""
        env = CRMQueryEnv()
        env.reset()
        
        action = {
            "tool": "search_customers",
            "arguments": {"tier": "Gold"}
        }
        
        # First query
        obs1, reward1, done1, info1 = env.step(action)
        base_reward = reward1.value
        
        # Repeated query
        obs2, reward2, done2, info2 = env.step(action)
        
        # Repeated query should have penalty
        assert "repeated_query" in reward2.components
        assert reward2.components["repeated_query"] < 0

    def test_different_queries_no_penalty(self):
        """Test that different queries don't trigger repeated penalty."""
        env = CRMQueryEnv()
        env.reset()
        
        action1 = {"tool": "search_customers", "arguments": {"tier": "Gold"}}
        action2 = {"tool": "search_customers", "arguments": {"tier": "Silver"}}
        
        obs1, reward1, done1, info1 = env.step(action1)
        obs2, reward2, done2, info2 = env.step(action2)
        
        # Second query is different, no repeated_query penalty
        assert "repeated_query" not in reward2.components or reward2.components.get("repeated_query") >= 0


class TestMemoryResetOnEpisode:
    """Test that memory resets properly between episodes."""

    def test_memory_reset_on_new_episode(self):
        """Test that memory clears on environment reset."""
        env = CRMQueryEnv()
        
        # First episode
        env.reset()
        action = {"tool": "search_customers", "arguments": {"tier": "Gold"}}
        env.step(action)
        
        entities_after_step = len(env.retrieved_entities['customers'])
        assert entities_after_step > 0
        
        # Reset environment
        env.reset()
        
        # Memory should be cleared
        assert len(env.retrieved_entities['customers']) == 0
        assert len(env.retrieved_entities['orders']) == 0
        assert len(env.retrieved_entities['tickets']) == 0
        assert len(env.step_summaries) == 0

    def test_query_history_reset(self):
        """Test that query history resets on environment reset."""
        env = CRMQueryEnv()
        
        env.reset()
        action = {"tool": "search_customers", "arguments": {"tier": "Gold"}}
        obs1, reward1, done1, info1 = env.step(action)
        
        # First repeat has penalty
        obs2, reward2, done2, info2 = env.step(action)
        assert reward2.components.get("repeated_query", 0) < 0
        
        # Reset
        env.reset()
        
        # After reset, query should not be in history
        obs3, reward3, done3, info3 = env.step(action)
        # No penalty for first occurrence in new episode
        assert reward3.components.get("repeated_query", 0) >= 0


class TestMemoryObservation:
    """Test memory information in observations."""

    def test_observation_includes_memory_cache(self):
        """Test that observations include memory cache."""
        env = CRMQueryEnv()
        obs = env.reset()
        
        assert hasattr(obs, 'memory_cache')
        assert isinstance(obs.memory_cache, dict)

    def test_observation_includes_step_summaries(self):
        """Test that observations include step summaries."""
        env = CRMQueryEnv()
        obs = env.reset()
        
        assert hasattr(obs, 'step_summaries')
        assert isinstance(obs.step_summaries, list)

    def test_memory_info_updated_in_observation(self):
        """Test that memory info is updated in subsequent observations."""
        env = CRMQueryEnv()
        obs1 = env.reset()
        
        initial_summaries = len(obs1.step_summaries)
        
        action = {"tool": "search_customers", "arguments": {"tier": "Gold"}}
        obs2, reward, done, info = env.step(action)
        
        assert len(obs2.step_summaries) > initial_summaries
        assert len(obs2.memory_cache['customers']) > 0