forked from BasedHardware/omi
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_retrieval_search.py
More file actions
249 lines (212 loc) · 9.79 KB
/
Copy pathtest_retrieval_search.py
File metadata and controls
249 lines (212 loc) · 9.79 KB
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
"""
Scenario: Retrieval/search critical path coverage.
These tests exercise real public retrieval/search APIs while replacing the
Pinecone/OpenAI boundary with a deterministic in-memory vector/search fake.
The fake is intentionally installed at the database.vector_db client seam:
routes, auth, Firestore persistence, vector upsert/delete calls, and result
hydration run through production code.
"""
from datetime import datetime, timedelta, timezone
from fakes.firestore import seed_conversation
from fakes.vector_search import install_vector_search_fakes
def _install_fakes(monkeypatch):
# Import after the client fixture has patched Firestore/Redis/Storage and
# imported the real app. Importing database.* at module collection time would
# instantiate real Google clients before the harness boundary is installed.
import database.vector_db as vector_db
fake_index, fake_embeddings = install_vector_search_fakes(monkeypatch, vector_db)
return vector_db, fake_index, fake_embeddings
def test_memory_search_reindexes_updates_and_delete_removes_result(client, auth_headers, fake_firestore, monkeypatch):
vector_db, fake_index, fake_embeddings = _install_fakes(monkeypatch)
from database.memory_vector_metadata import canonical_memory_provider_id
from utils.memory.canonical_required_processing import (
ProcessedRequiredMemory,
process_required_memory_item,
)
import utils.memory.short_term_promotion as promotion
from database.memory_outbox_worker import CanonicalMemoryOutboxSideEffects
real_side_effects = promotion._canonical_outbox_side_effects
def e2e_side_effects(*, db_client):
effects = real_side_effects(db_client=db_client)
return CanonicalMemoryOutboxSideEffects(
projection_upsert=lambda _item, _generation: True,
projection_delete=lambda _uid, _memory_id, _generation: True,
vector_upsert=effects.vector_upsert,
vector_delete=effects.vector_delete,
)
monkeypatch.setattr(promotion, "_canonical_outbox_side_effects", e2e_side_effects)
def process_and_drain(memory_id: str, run_id: str, *, process: bool = True) -> None:
if process:
processed = process_required_memory_item(
"123",
memory_id,
db_client=fake_firestore,
processor=lambda item: ProcessedRequiredMemory(content=item.content),
now=datetime.now(timezone.utc),
)
assert processed.processed is True
result = promotion._drain_canonical_outbox(
"123",
db_client=fake_firestore,
run_id=run_id,
now=datetime.now(timezone.utc) + timedelta(days=1),
)
assert result["errors"] == []
assert result["retryable_failure_count"] == 0
create = client.post(
"/v3/memories",
json={
"content": "David prefers canary deployments for backend rollouts",
"category": "system",
"visibility": "public",
},
headers=auth_headers,
)
assert create.status_code == 200, create.text
memory_id = create.json()["id"]
process_and_drain(memory_id, "e2e-retrieval-create")
search = client.post(
"/v1/tools/memories/search",
json={"query": "canary deployment preference", "limit": 5},
headers=auth_headers,
)
assert search.status_code == 200, search.text
assert "canary deployments" in search.json()["result_text"]
assert (
fake_embeddings.text_for_id(canonical_memory_provider_id("123", memory_id))
== "David prefers canary deployments for backend rollouts"
)
update = client.patch(
f"/v3/memories/{memory_id}",
params={"value": "David now prefers blue green releases for backend rollouts"},
headers=auth_headers,
)
assert update.status_code == 200, update.text
process_and_drain(memory_id, "e2e-retrieval-update")
old_search = client.post(
"/v1/tools/memories/search",
json={"query": "canary deployment preference", "limit": 5},
headers=auth_headers,
)
assert old_search.status_code == 200, old_search.text
assert "No memories found" in old_search.json()["result_text"]
new_search = client.post(
"/v1/tools/memories/search",
json={"query": "blue green backend releases", "limit": 5},
headers=auth_headers,
)
assert new_search.status_code == 200, new_search.text
assert "blue green releases" in new_search.json()["result_text"]
delete = client.delete(f"/v3/memories/{memory_id}", headers=auth_headers)
assert delete.status_code == 200, delete.text
process_and_drain(memory_id, "e2e-retrieval-delete", process=False)
after_delete = client.post(
"/v1/tools/memories/search",
json={"query": "blue green backend releases", "limit": 5},
headers=auth_headers,
)
assert after_delete.status_code == 200, after_delete.text
assert "No memories found" in after_delete.json()["result_text"]
assert fake_index.count(namespace=vector_db.MEMORIES_NAMESPACE) == 0
def test_action_item_search_tracks_create_update_and_delete(client, auth_headers, monkeypatch):
vector_db, fake_index, _ = _install_fakes(monkeypatch)
create = client.post(
"/v1/action-items",
json={"description": "Renew passport before Lisbon travel"},
headers=auth_headers,
)
assert create.status_code == 200, create.text
action_item_id = create.json()["id"]
search = client.get(
"/v1/action-items/search",
params={"query": "passport travel", "limit": 5},
headers=auth_headers,
)
assert search.status_code == 200, search.text
assert [item["id"] for item in search.json()["action_items"]] == [action_item_id]
update = client.patch(
f"/v1/action-items/{action_item_id}",
json={"description": "Book veterinarian appointment for Momo"},
headers=auth_headers,
)
assert update.status_code == 200, update.text
stale_search = client.get(
"/v1/action-items/search",
params={"query": "passport travel", "limit": 5},
headers=auth_headers,
)
assert stale_search.status_code == 200, stale_search.text
assert stale_search.json()["action_items"] == []
fresh_search = client.get(
"/v1/action-items/search",
params={"query": "veterinarian Momo", "limit": 5},
headers=auth_headers,
)
assert fresh_search.status_code == 200, fresh_search.text
assert [item["id"] for item in fresh_search.json()["action_items"]] == [action_item_id]
delete = client.delete(f"/v1/action-items/{action_item_id}", headers=auth_headers)
assert delete.status_code in (200, 204), delete.text
after_delete = client.get(
"/v1/action-items/search",
params={"query": "veterinarian Momo", "limit": 5},
headers=auth_headers,
)
assert after_delete.status_code == 200, after_delete.text
assert after_delete.json()["action_items"] == []
assert fake_index.count(namespace=vector_db.ACTION_ITEMS_NAMESPACE) == 0
def test_conversation_and_transcript_chunk_search_return_persisted_conversation_text(
client, auth_headers, sample_conversation_data, monkeypatch
):
vector_db, fake_index, _ = _install_fakes(monkeypatch)
from utils.conversations.transcript_chunks import build_transcript_chunks
conversation = dict(
sample_conversation_data,
id="conv-search-001",
structured={
**sample_conversation_data["structured"],
"title": "Infrastructure rollout",
"overview": "Discussed using feature flags for the production rollout.",
},
transcript_segments=[
{
"id": "seg-search-1",
"text": "The launch checklist says enable the Zagreb feature flag on Thursday.",
"speaker": "SPEAKER_00",
"is_user": True,
"start": 0.0,
"end": 3.0,
}
],
)
conversation["created_at"] = datetime.fromisoformat(conversation["created_at"].replace("Z", "+00:00"))
conversation["started_at"] = datetime.fromisoformat(conversation["started_at"].replace("Z", "+00:00"))
conversation["finished_at"] = datetime.fromisoformat(conversation["finished_at"].replace("Z", "+00:00"))
seed_conversation("123", conversation)
vector_db.upsert_vector2(
"123", conversation["id"], vector_db.embeddings.embed_query(conversation["structured"]["overview"]), {}
)
chunks = build_transcript_chunks(conversation["transcript_segments"], conversation["started_at"])
assert vector_db.upsert_transcript_chunk_vectors("123", conversation["id"], chunks) == 1
summary_search = client.post(
"/v1/tools/conversations/search",
json={"query": "production feature flags", "limit": 5, "include_transcript": False},
headers=auth_headers,
)
assert summary_search.status_code == 200, summary_search.text
assert "Infrastructure rollout" in summary_search.json()["result_text"]
chunk_search = client.post(
"/v1/tools/conversations/search-chunks",
json={"query": "Zagreb feature flag Thursday", "limit": 5},
headers=auth_headers,
)
assert chunk_search.status_code == 200, chunk_search.text
assert "Zagreb feature flag" in chunk_search.json()["result_text"]
vector_db.delete_transcript_chunk_vectors("123", conversation["id"])
no_chunk_search = client.post(
"/v1/tools/conversations/search-chunks",
json={"query": "Zagreb feature flag Thursday", "limit": 5},
headers=auth_headers,
)
assert no_chunk_search.status_code == 200, no_chunk_search.text
assert "No transcript excerpts found" in no_chunk_search.json()["result_text"]
assert fake_index.count(namespace=vector_db.TRANSCRIPT_CHUNKS_NAMESPACE) == 0