From 63f51385c413c8ae57826c817e9c5628cb849ede Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Nicol=C3=B2=20Boschi?= Date: Wed, 17 Dec 2025 13:17:38 +0100 Subject: [PATCH] fix: retain async fails (#40) * fix: retain async fails * fix: retain async fails --- hindsight-api/hindsight_api/api/http.py | 24 ++- .../hindsight_api/engine/memory_engine.py | 6 +- .../tests/test_http_api_integration.py | 182 ++++++++++++++++++ 3 files changed, 199 insertions(+), 13 deletions(-) diff --git a/hindsight-api/hindsight_api/api/http.py b/hindsight-api/hindsight_api/api/http.py index 39205963..d5b0cb54 100644 --- a/hindsight-api/hindsight_api/api/http.py +++ b/hindsight-api/hindsight_api/api/http.py @@ -1452,18 +1452,22 @@ def _register_routes(app: FastAPI): bank_id, ) + def parse_metadata(metadata): + """Parse result_metadata which may be a string or dict.""" + if metadata is None: + return {} + if isinstance(metadata, str): + return json.loads(metadata) + return metadata + return { "bank_id": bank_id, "operations": [ { "id": str(row["operation_id"]), "task_type": row["operation_type"], - "items_count": row["result_metadata"].get("items_count", 0) - if row["result_metadata"] - else 0, - "document_id": row["result_metadata"].get("document_id") - if row["result_metadata"] - else None, + "items_count": parse_metadata(row["result_metadata"]).get("items_count", 0), + "document_id": parse_metadata(row["result_metadata"]).get("document_id"), "created_at": row["created_at"].isoformat(), "status": row["status"], "error_message": row["error_message"], @@ -1499,7 +1503,7 @@ def _register_routes(app: FastAPI): async with acquire_with_retry(pool) as conn: # Check if operation exists and belongs to this memory bank result = await conn.fetchrow( - "SELECT bank_id FROM async_operations WHERE id = $1 AND bank_id = $2", op_uuid, bank_id + "SELECT bank_id FROM async_operations WHERE operation_id = $1 AND bank_id = $2", op_uuid, bank_id ) if not result: @@ -1508,7 +1512,7 @@ def _register_routes(app: FastAPI): ) # Delete the operation - await conn.execute("DELETE FROM async_operations WHERE id = $1", op_uuid) + await conn.execute("DELETE FROM async_operations WHERE operation_id = $1", op_uuid) return { "success": True, @@ -1769,13 +1773,13 @@ def _register_routes(app: FastAPI): async with acquire_with_retry(pool) as conn: await conn.execute( """ - INSERT INTO async_operations (id, bank_id, task_type, items_count) + INSERT INTO async_operations (operation_id, bank_id, operation_type, result_metadata) VALUES ($1, $2, $3, $4) """, operation_id, bank_id, "retain", - len(contents), + json.dumps({"items_count": len(contents)}), ) # Submit task to background queue diff --git a/hindsight-api/hindsight_api/engine/memory_engine.py b/hindsight-api/hindsight_api/engine/memory_engine.py index a0631ebb..65c9e483 100644 --- a/hindsight-api/hindsight_api/engine/memory_engine.py +++ b/hindsight-api/hindsight_api/engine/memory_engine.py @@ -311,7 +311,7 @@ class MemoryEngine: pool = await self._get_pool() async with acquire_with_retry(pool) as conn: result = await conn.fetchrow( - "SELECT id FROM async_operations WHERE id = $1", uuid.UUID(operation_id) + "SELECT operation_id FROM async_operations WHERE operation_id = $1", uuid.UUID(operation_id) ) if not result: # Operation was cancelled, skip processing @@ -369,7 +369,7 @@ class MemoryEngine: try: pool = await self._get_pool() async with acquire_with_retry(pool) as conn: - await conn.execute("DELETE FROM async_operations WHERE id = $1", uuid.UUID(operation_id)) + await conn.execute("DELETE FROM async_operations WHERE operation_id = $1", uuid.UUID(operation_id)) except Exception as e: logger.error(f"Failed to delete async operation record {operation_id}: {e}") @@ -386,7 +386,7 @@ class MemoryEngine: """ UPDATE async_operations SET status = 'failed', error_message = $2 - WHERE id = $1 + WHERE operation_id = $1 """, uuid.UUID(operation_id), truncated_error, diff --git a/hindsight-api/tests/test_http_api_integration.py b/hindsight-api/tests/test_http_api_integration.py index dcf3deb4..1707be5c 100644 --- a/hindsight-api/tests/test_http_api_integration.py +++ b/hindsight-api/tests/test_http_api_integration.py @@ -426,3 +426,185 @@ async def test_document_deletion(api_client): f"/v1/default/banks/{test_bank_id}/documents/sales-report-q1-2024" ) assert response.status_code == 404 + + +@pytest.mark.asyncio +async def test_async_retain(api_client): + """Test asynchronous retain functionality. + + When async=true is passed, the retain endpoint should: + 1. Return immediately with success and async_=true + 2. Process the content in the background + 3. Eventually store the memories + """ + import asyncio + + test_bank_id = f"async_retain_test_{datetime.now().timestamp()}" + + # Store memory with async=true + response = await api_client.post( + f"/v1/default/banks/{test_bank_id}/memories", + json={ + "async": True, + "items": [ + { + "content": "Alice is a senior engineer at TechCorp. She has been working on the authentication system for 5 years.", + "context": "team introduction" + } + ] + } + ) + assert response.status_code == 200 + result = response.json() + assert result["success"] is True + assert result["async"] is True, "Response should indicate async processing" + assert result["items_count"] == 1 + + # Check operations endpoint to see the pending operation + response = await api_client.get(f"/v1/default/banks/{test_bank_id}/operations") + assert response.status_code == 200 + ops_result = response.json() + assert "operations" in ops_result + + # Wait for async processing to complete (poll with timeout) + max_wait_seconds = 30 + poll_interval = 0.5 + elapsed = 0 + memories_found = False + + while elapsed < max_wait_seconds: + # Check if memories are stored + response = await api_client.get( + f"/v1/default/banks/{test_bank_id}/memories/list", + params={"limit": 10} + ) + assert response.status_code == 200 + items = response.json()["items"] + + if len(items) > 0: + memories_found = True + break + + await asyncio.sleep(poll_interval) + elapsed += poll_interval + + assert memories_found, f"Async retain did not complete within {max_wait_seconds} seconds" + + # Verify we can recall the stored memory + response = await api_client.post( + f"/v1/default/banks/{test_bank_id}/memories/recall", + json={ + "query": "Who works at TechCorp?", + "thinking_budget": 30 + } + ) + assert response.status_code == 200 + search_results = response.json() + assert len(search_results["results"]) > 0, "Should find the asynchronously stored memory" + + # Verify Alice is mentioned + found_alice = any("Alice" in r["text"] for r in search_results["results"]) + assert found_alice, "Should find Alice in search results" + + +@pytest.mark.asyncio +async def test_async_retain_parallel(api_client): + """Test multiple async retain operations running in parallel. + + Verifies that: + 1. Multiple async operations can be submitted concurrently + 2. All operations complete successfully + 3. The exact number of documents are processed + """ + import asyncio + + test_bank_id = f"async_parallel_test_{datetime.now().timestamp()}" + num_documents = 5 + + # Prepare multiple documents to retain + documents = [ + { + "content": f"Document {i}: This is test content about Person{i} who works at Company{i}.", + "context": f"test document {i}", + "document_id": f"doc_{i}" + } + for i in range(num_documents) + ] + + # Submit all async retain operations in parallel + async def submit_async_retain(doc): + return await api_client.post( + f"/v1/default/banks/{test_bank_id}/memories", + json={ + "async": True, + "items": [doc] + } + ) + + # Run all submissions concurrently + responses = await asyncio.gather(*[submit_async_retain(doc) for doc in documents]) + + # Verify all submissions succeeded + for i, response in enumerate(responses): + assert response.status_code == 200, f"Document {i} submission failed" + result = response.json() + assert result["success"] is True + assert result["async"] is True + + # Check operations endpoint - should show pending operations + response = await api_client.get(f"/v1/default/banks/{test_bank_id}/operations") + assert response.status_code == 200 + + # Wait for all async operations to complete (poll with timeout) + max_wait_seconds = 60 + poll_interval = 1.0 + elapsed = 0 + all_docs_processed = False + + while elapsed < max_wait_seconds: + # Check document count + response = await api_client.get(f"/v1/default/banks/{test_bank_id}/documents") + assert response.status_code == 200 + docs = response.json()["items"] + + if len(docs) >= num_documents: + all_docs_processed = True + break + + await asyncio.sleep(poll_interval) + elapsed += poll_interval + + assert all_docs_processed, f"Expected {num_documents} documents, but only {len(docs)} were processed within {max_wait_seconds} seconds" + + # Verify exact document count + response = await api_client.get(f"/v1/default/banks/{test_bank_id}/documents") + assert response.status_code == 200 + final_docs = response.json()["items"] + assert len(final_docs) == num_documents, f"Expected exactly {num_documents} documents, got {len(final_docs)}" + + # Verify each document exists + doc_ids = {doc["id"] for doc in final_docs} + for i in range(num_documents): + assert f"doc_{i}" in doc_ids, f"Document doc_{i} not found" + + # Verify memories were created for all documents + response = await api_client.get( + f"/v1/default/banks/{test_bank_id}/memories/list", + params={"limit": 100} + ) + assert response.status_code == 200 + memories = response.json()["items"] + assert len(memories) >= num_documents, f"Expected at least {num_documents} memories, got {len(memories)}" + + # Verify we can recall content from different documents + for i in [0, num_documents - 1]: # Check first and last + response = await api_client.post( + f"/v1/default/banks/{test_bank_id}/memories/recall", + json={ + "query": f"Who works at Company{i}?", + "thinking_budget": 30 + } + ) + assert response.status_code == 200 + results = response.json()["results"] + assert len(results) > 0, f"Should find memories for document {i}"