|
18 | 18 | import app |
19 | 19 | import automem.api.recall as recall_module |
20 | 20 | import automem.search.runtime_recall_helpers as recall_helpers |
| 21 | +from automem.utils.scoring import _compute_metadata_score |
21 | 22 | from tests.support.fake_graph import FakeGraph |
22 | 23 |
|
23 | 24 |
|
@@ -138,6 +139,74 @@ def test_tag_scoped_vector_fetch_respects_cap(overfetch_env, monkeypatch): |
138 | 139 | assert overfetch_env.last_limit == 20 |
139 | 140 |
|
140 | 141 |
|
| 142 | +def test_context_tags_widen_vector_fetch_before_scoring(monkeypatch) -> None: |
| 143 | + monkeypatch.setattr(recall_module, "RECALL_VECTOR_OVERFETCH", 4) |
| 144 | + monkeypatch.setattr(recall_module, "RECALL_VECTOR_FETCH_CAP", 200) |
| 145 | + |
| 146 | + requested_limits = [] |
| 147 | + |
| 148 | + def _result(memory_id, score, *, tags=None): |
| 149 | + return { |
| 150 | + "id": memory_id, |
| 151 | + "score": score, |
| 152 | + "match_score": score, |
| 153 | + "match_type": "vector", |
| 154 | + "source": "qdrant", |
| 155 | + "memory": { |
| 156 | + "id": memory_id, |
| 157 | + "content": f"unrelated content for {memory_id}", |
| 158 | + "tags": tags or [], |
| 159 | + "importance": 0.0, |
| 160 | + "confidence": 0.0, |
| 161 | + "type": "Context", |
| 162 | + "enriched": True, |
| 163 | + "timestamp": "2026-05-18T00:00:00+00:00", |
| 164 | + }, |
| 165 | + "relations": [], |
| 166 | + } |
| 167 | + |
| 168 | + vector_results = [_result(f"decoy-{i}", 0.95 - i * 0.001) for i in range(24)] + [ |
| 169 | + _result("context-target", 0.01, tags=["target-tag"]), |
| 170 | + ] |
| 171 | + |
| 172 | + def _vector_search(_qdrant, _graph, _query, _embedding, limit, seen_ids, *_args): |
| 173 | + requested_limits.append(limit) |
| 174 | + matches = [dict(result) for result in vector_results[:limit]] |
| 175 | + for result in matches: |
| 176 | + seen_ids.add(result["id"]) |
| 177 | + return matches |
| 178 | + |
| 179 | + with app.app.test_request_context( |
| 180 | + "/recall?query=quarterly%20metrics&context_tags=target-tag&limit=5¤t_only=false" |
| 181 | + ): |
| 182 | + response = recall_module.handle_recall( |
| 183 | + get_memory_graph=lambda: None, |
| 184 | + get_qdrant_client=lambda: object(), |
| 185 | + normalize_tag_list=lambda value: value if isinstance(value, list) else [], |
| 186 | + normalize_timestamp=lambda value: value, |
| 187 | + parse_time_expression=lambda _value: (None, None), |
| 188 | + extract_keywords=lambda _query: ["quarterly", "metrics"], |
| 189 | + compute_metadata_score=_compute_metadata_score, |
| 190 | + result_passes_filters=lambda *_args, **_kwargs: True, |
| 191 | + graph_keyword_search=lambda *_args, **_kwargs: [], |
| 192 | + vector_search=_vector_search, |
| 193 | + vector_filter_only_tag_search=lambda *_args, **_kwargs: [], |
| 194 | + metadata_keyword_search=lambda *_args, **_kwargs: [], |
| 195 | + recall_max_limit=50, |
| 196 | + logger=SimpleNamespace( |
| 197 | + debug=lambda *_args, **_kwargs: None, |
| 198 | + info=lambda *_args, **_kwargs: None, |
| 199 | + exception=lambda *_args, **_kwargs: None, |
| 200 | + ), |
| 201 | + ) |
| 202 | + |
| 203 | + data = response.get_json() |
| 204 | + ids = [result["id"] for result in data["results"]] |
| 205 | + assert requested_limits == [50] |
| 206 | + assert "context-target" in ids |
| 207 | + assert len(ids) <= 5 |
| 208 | + |
| 209 | + |
141 | 210 | def test_vector_overfetch_hydrates_relations_after_trim(overfetch_env, monkeypatch): |
142 | 211 | relation_calls = [] |
143 | 212 |
|
|
0 commit comments