Compare commits

...
2 changed files with 109 additions and 35 deletions
@@ -156,13 +156,22 @@ async def retain_batch(
)
if not extracted_facts:
# Still need to create document if document_id was provided
# Still need to create document if document_id was provided or chunks exist
from collections import defaultdict
docs_tracked = 0
async with acquire_with_retry(pool) as conn:
async with conn.transaction():
await fact_storage.ensure_bank_exists(conn, bank_id)
# Handle document tracking even with no facts
# Group contents by document_id (consistent with normal path)
contents_by_doc_early = defaultdict(list)
for idx, content_dict in enumerate(contents_dicts):
doc_id = content_dict.get("document_id")
contents_by_doc_early[doc_id].append((idx, content_dict))
if document_id:
# Legacy: single document_id parameter
combined_content = "\n".join([c.get("content", "") for c in contents_dicts])
# Collect tags from all content items and merge with document_tags
all_tags = set(document_tags or [])
@@ -187,45 +196,57 @@ async def retain_batch(
await fact_storage.handle_document_tracking(
conn, bank_id, document_id, combined_content, is_first_batch, retain_params, merged_tags
)
docs_tracked += 1
else:
# Check for per-item document_ids
from collections import defaultdict
# Handle per-item document_ids and/or chunks (mirrors normal path logic)
has_any_doc_ids = any(item.get("document_id") for item in contents_dicts)
contents_by_doc = defaultdict(list)
for idx, content_dict in enumerate(contents_dicts):
doc_id = content_dict.get("document_id")
if doc_id:
contents_by_doc[doc_id].append((idx, content_dict))
if has_any_doc_ids or chunks:
for original_doc_id, doc_contents in contents_by_doc_early.items():
should_create_doc = (original_doc_id is not None) or chunks
if not should_create_doc:
continue
for doc_id, doc_contents in contents_by_doc.items():
combined_content = "\n".join([c.get("content", "") for _, c in doc_contents])
# Collect tags from all content items for this document and merge with document_tags
all_tags = set(document_tags or [])
for _, item in doc_contents:
item_tags = item.get("tags", []) or []
all_tags.update(item_tags)
merged_tags = list(all_tags)
actual_doc_id = original_doc_id
if actual_doc_id is None:
# No document_id but have chunks - generate one
actual_doc_id = str(uuid.uuid4())
retain_params = {}
if doc_contents:
first_item = doc_contents[0][1]
if first_item.get("context"):
retain_params["context"] = first_item["context"]
if first_item.get("event_date"):
retain_params["event_date"] = (
first_item["event_date"].isoformat()
if hasattr(first_item["event_date"], "isoformat")
else str(first_item["event_date"])
)
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn, bank_id, doc_id, combined_content, is_first_batch, retain_params, merged_tags
)
combined_content = "\n".join([c.get("content", "") for _, c in doc_contents])
all_tags = set(document_tags or [])
for _, item in doc_contents:
item_tags = item.get("tags", []) or []
all_tags.update(item_tags)
merged_tags = list(all_tags)
retain_params = {}
if doc_contents:
first_item = doc_contents[0][1]
if first_item.get("context"):
retain_params["context"] = first_item["context"]
if first_item.get("event_date"):
retain_params["event_date"] = (
first_item["event_date"].isoformat()
if hasattr(first_item["event_date"], "isoformat")
else str(first_item["event_date"])
)
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
await fact_storage.handle_document_tracking(
conn,
bank_id,
actual_doc_id,
combined_content,
is_first_batch,
retain_params,
merged_tags,
)
docs_tracked += 1
total_time = time.time() - start_time
doc_status = f"{docs_tracked} document(s) tracked" if docs_tracked > 0 else "no document tracked"
logger.info(
f"RETAIN_BATCH COMPLETE: 0 facts extracted from {len(contents)} contents in {total_time:.3f}s (document tracked, no facts)"
f"RETAIN_BATCH COMPLETE: 0 facts extracted from {len(contents)} contents in {total_time:.3f}s ({doc_status}, no facts)"
)
return [[] for _ in contents], usage
+54 -1
View File
@@ -2,9 +2,13 @@
Tests for document tracking and upsert functionality.
"""
import logging
import pytest
from datetime import datetime, timezone
from unittest.mock import patch
import pytest
from hindsight_api import RequestContext
from hindsight_api.engine.response_models import TokenUsage
@pytest.mark.asyncio
@@ -311,3 +315,52 @@ async def test_document_persisted_with_zero_facts_async_submit(memory, request_c
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_document_stored_without_chunks_when_zero_facts(memory_no_llm_verify, request_context):
"""
Regression test: when 0 facts are extracted from chunked content, the document row
must be stored but no chunk rows should be written.
"""
bank_id = f"test_zero_facts_no_chunks_{datetime.now(timezone.utc).timestamp()}"
document_id = "doc-zero-facts-chunked"
# Content large enough to exceed default retain_chunk_size (3000 chars) so chunking is triggered
content = "Alice works at Google. " * 200 # ~4600 chars
async def mock_llm_zero_facts(*args, **kwargs):
response = {"facts": []}
if kwargs.get("return_usage", False):
return response, TokenUsage(input_tokens=10, output_tokens=2)
return response
try:
with patch("hindsight_api.engine.llm_wrapper.LLMProvider.call", new=mock_llm_zero_facts):
units = await memory_no_llm_verify.retain_async(
bank_id=bank_id,
content=content,
document_id=document_id,
request_context=request_context,
)
assert units == [], "Should return no memory units when LLM extracts zero facts"
# Document row must exist
doc = await memory_no_llm_verify.get_document(document_id, bank_id, request_context=request_context)
assert doc is not None, "Document row must be stored even when zero facts are extracted"
assert doc["id"] == document_id
assert doc["memory_unit_count"] == 0
# No chunk rows should be stored
pool = await memory_no_llm_verify._get_pool()
async with pool.acquire() as conn:
chunk_count = await conn.fetchval(
"SELECT COUNT(*) FROM chunks WHERE document_id = $1 AND bank_id = $2",
document_id,
bank_id,
)
assert chunk_count == 0, "No chunk rows should be stored when zero facts are extracted"
finally:
await memory_no_llm_verify.delete_bank(bank_id, request_context=request_context)