44
55from api .services import extraction , proxy
66from api .services .proxy import ProviderResponse
7- from api .services .users import regenerate_user_key , update_user_provider_config
7+ from api .services .users import regenerate_user_key , update_user_provider_config
88
99
1010@pytest .mark .asyncio
1111async def test_proxy_uses_user_retrieval_config (monkeypatch ) -> None :
1212 captured : dict [str , object ] = {}
1313
14- async def fake_retrieve_memories (user_id , query , db , limit = None , threshold = None ):
14+ async def fake_retrieve_memories (user_id , query , db , limit = None , threshold = None , namespace = "default" ):
1515 captured ["user_id" ] = user_id
1616 captured ["query" ] = query
1717 captured ["limit" ] = limit
1818 captured ["threshold" ] = threshold
19+ captured ["namespace" ] = namespace
1920 return []
2021
2122 async def fake_log_retrieval (user_id , conversation_id , query , memories , db ):
@@ -74,23 +75,26 @@ async def fetchrow(self, query, *args):
7475async def test_extracted_memory_storage_uses_user_dedup_threshold (monkeypatch ) -> None :
7576 captured : dict [str , object ] = {}
7677
77- async def fake_store_memory_with_deduplication (
78- user_id ,
79- content ,
80- embedding ,
81- conversation_id ,
82- confidence ,
83- db ,
84- dedup_threshold = None ,
85- status = "approved" ,
86- category = "general" ,
87- source = "manual" ,
88- ):
89- captured ["dedup_threshold" ] = dedup_threshold
90- captured ["status" ] = status
91- captured ["category" ] = category
92- captured ["source" ] = source
93- return {"action" : "inserted" , "memory" : {"id" : "memory-1" }}
78+ async def fake_store_memory_with_deduplication (
79+ user_id ,
80+ content ,
81+ embedding ,
82+ conversation_id ,
83+ confidence ,
84+ db ,
85+ dedup_threshold = None ,
86+ status = "approved" ,
87+ category = "general" ,
88+ source = "manual" ,
89+ namespace = "default" ,
90+ ):
91+ captured ["dedup_threshold" ] = dedup_threshold
92+ captured ["status" ] = status
93+ captured ["category" ] = category
94+ captured ["source" ] = source
95+ captured ["namespace" ] = namespace
96+ from uuid import uuid4 as _uuid4
97+ return {"action" : "inserted" , "memory" : {"id" : _uuid4 ()}}
9498
9599 def fake_embed_batch (texts ):
96100 captured ["texts" ] = texts
@@ -99,24 +103,24 @@ def fake_embed_batch(texts):
99103 monkeypatch .setattr (extraction , "store_memory_with_deduplication" , fake_store_memory_with_deduplication )
100104 monkeypatch .setattr ("api.services.embedding.embed_batch" , fake_embed_batch )
101105
102- stored_count = await extraction .store_extracted_memories (
106+ stored_count , _stored_refs = await extraction .store_extracted_memories (
103107 uuid4 (),
104108 uuid4 (),
105109 ["User prefers FastAPI" ],
106110 object (),
107111 0.61 ,
108112 )
109113
110- assert stored_count == 1
111- assert captured ["texts" ] == ["User prefers FastAPI" ]
112- assert captured ["dedup_threshold" ] == 0.61
113- assert captured ["status" ] == "pending"
114- assert captured ["category" ] == "preferences"
115- assert captured ["source" ] == "extraction"
114+ assert stored_count == 1
115+ assert captured ["texts" ] == ["User prefers FastAPI" ]
116+ assert captured ["dedup_threshold" ] == 0.61
117+ assert captured ["status" ] == "pending"
118+ assert captured ["category" ] == "preferences"
119+ assert captured ["source" ] == "extraction"
116120
117121
118122@pytest .mark .asyncio
119- async def test_regenerate_user_key_removes_all_old_issued_keys () -> None :
123+ async def test_regenerate_user_key_removes_all_old_issued_keys () -> None :
120124 class FakeTransaction :
121125 async def __aenter__ (self ):
122126 return self
@@ -160,73 +164,73 @@ async def execute(self, query, *args):
160164
161165 await regenerate_user_key (user , db )
162166
163- assert db .deleted_user_id == user ["id" ]
164- assert db .deleted_with_hash_filter is False
165- assert db .inserted_key_name == "default"
166-
167-
168- @pytest .mark .asyncio
169- async def test_update_user_provider_config_persists_extraction_model () -> None :
170- class FakeDb :
171- def __init__ (self ) -> None :
172- self .query = ""
173- self .args : tuple [object , ...] = ()
174-
175- async def fetchrow (self , query , * args ):
176- self .query = query
177- self .args = args
178- return {
179- "id" : args [0 ],
180- "external_id" : "external-user" ,
181- "created_at" : None ,
182- "max_memories_injected" : 5 ,
183- "retrieval_threshold" : 0.5 ,
184- "dedup_threshold" : 0.95 ,
185- "extraction_provider" : args [1 ],
186- "extraction_model" : args [2 ],
187- "openai_api_key_encrypted" : None ,
188- "gemini_api_key_encrypted" : None ,
189- "anthropic_api_key_encrypted" : None ,
190- }
191-
192- db = FakeDb ()
193- user_id = uuid4 ()
194-
195- response = await update_user_provider_config (
196- user_id ,
197- "gemini" ,
198- "gemini-1.5-flash" ,
199- None ,
200- None ,
201- None ,
202- False ,
203- False ,
204- False ,
205- db ,
206- )
207-
208- assert "extraction_model" in db .query
209- assert db .args == (user_id , "gemini" , "gemini-1.5-flash" )
210- assert response ["extraction_provider" ] == "gemini"
211- assert response ["extraction_model" ] == "gemini-1.5-flash"
212-
213-
214- @pytest .mark .asyncio
215- async def test_update_user_provider_config_rejects_blank_extraction_model () -> None :
216- class FakeDb :
217- async def fetchrow (self , query , * args ):
218- raise AssertionError ("database should not be called for invalid model" )
219-
220- with pytest .raises (ValueError , match = "Extraction model is required" ):
221- await update_user_provider_config (
222- uuid4 (),
223- "openai" ,
224- " " ,
225- None ,
226- None ,
227- None ,
228- False ,
229- False ,
230- False ,
231- FakeDb (),
232- )
167+ assert db .deleted_user_id == user ["id" ]
168+ assert db .deleted_with_hash_filter is False
169+ assert db .inserted_key_name == "default"
170+
171+
172+ @pytest .mark .asyncio
173+ async def test_update_user_provider_config_persists_extraction_model () -> None :
174+ class FakeDb :
175+ def __init__ (self ) -> None :
176+ self .query = ""
177+ self .args : tuple [object , ...] = ()
178+
179+ async def fetchrow (self , query , * args ):
180+ self .query = query
181+ self .args = args
182+ return {
183+ "id" : args [0 ],
184+ "external_id" : "external-user" ,
185+ "created_at" : None ,
186+ "max_memories_injected" : 5 ,
187+ "retrieval_threshold" : 0.5 ,
188+ "dedup_threshold" : 0.95 ,
189+ "extraction_provider" : args [1 ],
190+ "extraction_model" : args [2 ],
191+ "openai_api_key_encrypted" : None ,
192+ "gemini_api_key_encrypted" : None ,
193+ "anthropic_api_key_encrypted" : None ,
194+ }
195+
196+ db = FakeDb ()
197+ user_id = uuid4 ()
198+
199+ response = await update_user_provider_config (
200+ user_id ,
201+ "gemini" ,
202+ "gemini-1.5-flash" ,
203+ None ,
204+ None ,
205+ None ,
206+ False ,
207+ False ,
208+ False ,
209+ db ,
210+ )
211+
212+ assert "extraction_model" in db .query
213+ assert db .args == (user_id , "gemini" , "gemini-1.5-flash" )
214+ assert response ["extraction_provider" ] == "gemini"
215+ assert response ["extraction_model" ] == "gemini-1.5-flash"
216+
217+
218+ @pytest .mark .asyncio
219+ async def test_update_user_provider_config_rejects_blank_extraction_model () -> None :
220+ class FakeDb :
221+ async def fetchrow (self , query , * args ):
222+ raise AssertionError ("database should not be called for invalid model" )
223+
224+ with pytest .raises (ValueError , match = "Extraction model is required" ):
225+ await update_user_provider_config (
226+ uuid4 (),
227+ "openai" ,
228+ " " ,
229+ None ,
230+ None ,
231+ None ,
232+ False ,
233+ False ,
234+ False ,
235+ FakeDb (),
236+ )
0 commit comments