From b4b5c44a8702bd55d9b17126bf73dea38512ba25 Mon Sep 17 00:00:00 2001 From: abix5 <6427182+abix5@users.noreply.github.com> Date: Mon, 16 Feb 2026 12:04:48 +0300 Subject: [PATCH] fix: propagate document tags in async retain path (#374) --- .../hindsight_api/engine/interface.py | 6 +- .../hindsight_api/engine/memory_engine.py | 8 ++- hindsight-api/tests/test_async_retain_tags.py | 70 +++++++++++++++++++ 3 files changed, 82 insertions(+), 2 deletions(-) create mode 100644 hindsight-api/tests/test_async_retain_tags.py diff --git a/hindsight-api/hindsight_api/engine/interface.py b/hindsight-api/hindsight_api/engine/interface.py index ed41e808..6115dfde 100644 --- a/hindsight-api/hindsight_api/engine/interface.py +++ b/hindsight-api/hindsight_api/engine/interface.py @@ -48,6 +48,7 @@ class MemoryEngineInterface(ABC): contents: list[dict[str, Any]], *, request_context: "RequestContext", + document_tags: list[str] | None = None, ) -> dict[str, Any]: """ Retain a batch of memory items. @@ -55,8 +56,9 @@ class MemoryEngineInterface(ABC): Args: bank_id: The memory bank ID. contents: List of content dicts with 'content', optional 'event_date', - 'context', 'metadata', 'document_id'. + 'context', 'metadata', 'document_id', and per-item 'tags'. request_context: Request context for authentication. + document_tags: Optional tags applied to all items in the batch. Returns: Dict with processing results. @@ -561,6 +563,7 @@ class MemoryEngineInterface(ABC): contents: list[dict[str, Any]], *, request_context: "RequestContext", + document_tags: list[str] | None = None, ) -> dict[str, Any]: """ Submit a batch retain operation to run asynchronously. @@ -569,6 +572,7 @@ class MemoryEngineInterface(ABC): bank_id: The memory bank ID. contents: List of content dicts to retain. request_context: Request context for authentication. + document_tags: Optional tags applied to all items in the async batch. Returns: Dict with operation_id and items_count. diff --git a/hindsight-api/hindsight_api/engine/memory_engine.py b/hindsight-api/hindsight_api/engine/memory_engine.py index 4438d777..a45e9613 100644 --- a/hindsight-api/hindsight_api/engine/memory_engine.py +++ b/hindsight-api/hindsight_api/engine/memory_engine.py @@ -540,6 +540,7 @@ class MemoryEngine(MemoryEngineInterface): if not bank_id: raise ValueError("bank_id is required for batch retain task") contents = task_dict.get("contents", []) + document_tags = task_dict.get("document_tags") logger.info( f"[BATCH_RETAIN_TASK] Starting background batch retain for bank_id={bank_id}, {len(contents)} items" @@ -557,7 +558,12 @@ class MemoryEngine(MemoryEngineInterface): tenant_id=task_dict.get("_tenant_id"), api_key_id=task_dict.get("_api_key_id"), ) - await self.retain_batch_async(bank_id=bank_id, contents=contents, request_context=context) + await self.retain_batch_async( + bank_id=bank_id, + contents=contents, + document_tags=document_tags, + request_context=context, + ) logger.info(f"[BATCH_RETAIN_TASK] Completed background batch retain for bank_id={bank_id}") diff --git a/hindsight-api/tests/test_async_retain_tags.py b/hindsight-api/tests/test_async_retain_tags.py new file mode 100644 index 00000000..24fd4380 --- /dev/null +++ b/hindsight-api/tests/test_async_retain_tags.py @@ -0,0 +1,70 @@ +"""Unit tests for async retain tag propagation.""" + +from unittest.mock import AsyncMock + +import pytest + +from hindsight_api.engine.memory_engine import MemoryEngine +from hindsight_api.models import RequestContext + + +@pytest.mark.asyncio +async def test_submit_async_retain_includes_document_tags_in_task_payload(): + """submit_async_retain should include document_tags in queued task payload.""" + engine = MemoryEngine.__new__(MemoryEngine) + engine._authenticate_tenant = AsyncMock() + engine._submit_async_operation = AsyncMock(return_value={"operation_id": "op-1"}) + + request_context = RequestContext(tenant_id="tenant-a", api_key_id="key-a") + contents = [{"content": "Async retain payload test."}] + document_tags = ["scope:tools", "user:alice"] + + result = await MemoryEngine.submit_async_retain( + engine, + bank_id="bank-1", + contents=contents, + document_tags=document_tags, + request_context=request_context, + ) + + assert result == {"operation_id": "op-1", "items_count": 1} + engine._authenticate_tenant.assert_awaited_once_with(request_context) + engine._submit_async_operation.assert_awaited_once() + + kwargs = engine._submit_async_operation.await_args.kwargs + assert kwargs["bank_id"] == "bank-1" + assert kwargs["operation_type"] == "retain" + assert kwargs["task_type"] == "batch_retain" + assert kwargs["task_payload"]["contents"] == contents + assert kwargs["task_payload"]["document_tags"] == document_tags + assert kwargs["task_payload"]["_tenant_id"] == "tenant-a" + assert kwargs["task_payload"]["_api_key_id"] == "key-a" + + +@pytest.mark.asyncio +async def test_handle_batch_retain_forwards_document_tags_to_retain_batch_async(): + """Worker handler should forward document_tags from task payload.""" + engine = MemoryEngine.__new__(MemoryEngine) + engine.retain_batch_async = AsyncMock(return_value={"items_count": 1}) + + task_dict = { + "bank_id": "bank-1", + "contents": [{"content": "Forward tags test."}], + "document_tags": ["scope:client"], + "_tenant_id": "tenant-a", + "_api_key_id": "key-a", + } + + await MemoryEngine._handle_batch_retain(engine, task_dict) + + engine.retain_batch_async.assert_awaited_once() + kwargs = engine.retain_batch_async.await_args.kwargs + assert kwargs["bank_id"] == "bank-1" + assert kwargs["contents"] == task_dict["contents"] + assert kwargs["document_tags"] == ["scope:client"] + + request_context = kwargs["request_context"] + assert request_context.internal is True + assert request_context.user_initiated is True + assert request_context.tenant_id == "tenant-a" + assert request_context.api_key_id == "key-a"