* perf: 3-phase retain pipeline — fix deadlocks, cap temporal links, query-time entity expansion
Major retain pipeline overhaul addressing deadlocks, write amplification,
and TimeoutErrors. Restructures retain into three phases:
Phase 1: Entity resolution on separate connection (read-heavy)
Phase 2: Core write transaction (atomic) — facts, unit_entities, links
Phase 3: Best-effort display data (error-isolated) — entity viz links, stats
Key changes:
- Sorted bulk INSERT FROM unnest() prevents deadlocks
- Temporal links capped to top-20 per unit (95% reduction)
- Batched semantic ANN via temp table + LATERAL
- Query-time entity expansion via unit_entities self-join
- Entity viz links moved to Phase 3 (post-transaction)
- HINDSIGHT_API_RETAIN_MAX_CONCURRENT config (default: 32)
* fix: increase semantic link top_k from 5 to 20
The hardcoded top_k=5 was artificially limiting semantic link creation.
Link expansion retrieval can consume up to budget (50-200) semantic
neighbors per seed set, but each fact only had 5 outgoing edges — making
the bidirectional graph very sparse.
Increasing to 20 gives retrieval 4x more edges to work with. The ANN
probe cost is unchanged (same HNSW traversal per fact, just returning
more rows). INSERT cost is negligible (~14k rows via bulk INSERT).
Also: all 18 TimeoutErrors in the latest benchmark (beam-1m-u20) were
from Gemini LLM calls, zero from the database — confirming the entity
resolution split eliminated DB timeouts entirely.
* perf: move semantic ANN search to Phase 1 to avoid transaction timeouts
The batched LATERAL ANN query (700 HNSW probes) was the last remaining
source of DB TimeoutErrors — all 29 in the latest benchmark were from
create_semantic_links_batch inside the Phase 2 write transaction.
Split semantic link creation into three phases:
- Phase 1 (separate conn, autocommit): ANN search via temp table + LATERAL.
No transaction locks, no contention with concurrent writers.
- Phase 2 (write transaction): within-batch numpy similarities (instant) +
INSERT of both within-batch and Phase 1 ANN results. No DB reads.
- Phase 3 (flush_pending_stats): future hook point for re-checking ANN
results after commit to catch links missed by concurrent batches.
Also adds 7 unit tests for compute_semantic_links_within_batch covering
empty input, identical/orthogonal embeddings, threshold filtering, top_k
cap, and tuple structure validation.
* fix: handle placeholder unit_ids in Phase 1 ANN search (not valid UUIDs)
* test: add Phase 1 ANN cross-batch test + configurable test PG port
- New test_semantic_links_phase1_ann_cross_batch verifies that the Phase 1
ANN search with placeholder unit IDs correctly creates cross-batch
semantic links after remapping to real IDs.
- Test PG port now configurable via HINDSIGHT_TEST_PG_PORT env var
(default: 5556) to avoid conflicts with running benchmark daemons.
* perf: remove retry_with_backoff from retain, set semaphore default to 4
Remove retry_with_backoff from _run_db_work and _run_delta_db_work:
- Deadlocks are prevented by sorted bulk INSERT (no need for retry)
- Transient timeouts are handled by the worker poller's task-level retry
(3 attempts, 60s spacing) which is better than rapid internal retries
that amplify I/O pressure during contention storms
Set HINDSIGHT_API_RETAIN_MAX_CONCURRENT default from 32 to 4:
- The semaphore gates Phase 1 (ANN + entity resolution) + Phase 2 (writes)
- At 4 concurrent, HNSW index I/O is manageable; at 10+ concurrent the
probes saturate disk and cause cascading timeouts
- LLM extraction still runs at full parallelism (semaphore acquired after)
* fix: add fact_type filter to Phase 1 ANN query to use per-bank HNSW indexes
The LATERAL ANN query was falling back to sequential scan + sort (90ms/probe)
because the per-bank HNSW indexes are partial indexes filtered on fact_type.
Without fact_type in the WHERE clause, PostgreSQL couldn't use them.
Fix: iterate over ('world', 'experience') and run one HNSW-indexed ANN per
type. EXPLAIN shows 8ms/probe (was 90ms) — 11x faster.
700 probes × 8ms × 2 types = ~11s total (was ~63s via seq scan).
* fix: scope temporal links by fact_type + add integration tests
Temporal links now filter by fact_type in the LATERAL query — world facts
only link to world facts, experience to experience. This matches how
retrieval filters results and avoids wasted cross-type link rows.
New integration tests:
- test_semantic_ann_uses_hnsw_index: verifies Phase 1 ANN creates
cross-batch semantic links (tests fact_type filter + placeholder remap)
- test_temporal_links_scoped_by_fact_type: verifies world facts get
temporal links to other world facts but NOT to experience facts
* fix: tolerate individual chunk LLM failures instead of failing entire batch
Changed asyncio.gather(*tasks) to asyncio.gather(*tasks, return_exceptions=True)
in both chunk-level and content-level fact extraction. A single chunk timeout
(e.g., Gemini >90s) no longer discards all other successfully extracted facts.
For a 50MB document with 17k chunks, even a 2% chunk failure rate previously
caused 0 completions (entire batch discarded). Now 16,700 facts are extracted
and only the 300 failed chunks are skipped with a warning log.
* fix: batch temporal LATERAL query for large documents (16k+ chunks)
The LATERAL query for temporal links passed all unit_ids at once into
unnest(), causing PostgreSQL timeouts on documents with 16k+ chunks.
Split into batches of 500 units per query to keep each under the
command_timeout.
Also identified: HNSW index creation on shared pg0 instances with
50k+ existing units exceeds the 60s command_timeout. This is a
test infrastructure issue (shared pg0 accumulates data) but also
affects production when creating new banks on large instances.
* feat: streaming chunk batching for large documents (RETAIN_CHUNK_BATCH_SIZE)
Process chunks in mini-batches of N (default 500), committing each batch
to the DB before starting the next. This prevents OOM kills on large
documents (50MB / 17k+ chunks) by keeping only ~500 facts + embeddings
in memory at a time instead of 50k+.
Each mini-batch goes through the full Phase 1 → 2 → 3 pipeline
independently, sharing the same document_id. On recovery (process dies
mid-way), delta retain detects already-committed chunks via content_hash
and skips them — only remaining chunks get re-extracted.
Config: HINDSIGHT_API_RETAIN_CHUNK_BATCH_SIZE (default: 500, 0 to disable)
Per-bank configurable via the hierarchical config system.
Tests:
- test_streaming_chunk_batching_produces_same_facts
- test_streaming_chunk_batching_recovery (delta retain skips committed chunks)
- test_streaming_disabled_for_small_docs
* perf(retain): producer-consumer pipeline + deferred semantic ANN
Replace the sequential streaming loop with a producer-consumer pipeline:
- LLM producer fires concurrent chunk extractions (semaphore-bounded)
- DB consumer drains queue in batches, runs Phase 1+2+3 per batch
- LLM and DB work overlap instead of running sequentially
Defer semantic links to a single final ANN pass after all batches commit:
- Remove within-batch semantic links from Phase 2 (was 2.6s/batch)
- Run parallel ANN (4 connections) after all facts committed
- top_k reduced from 50 to 20 (recall uses at most 20 neighbors)
- Recovery via operation result_metadata checkpoint
Additional optimizations:
- skip_exists_check on temporal/causal link INSERT (saves ~0.5s/batch)
- WHERE EXISTS guard on semantic link INSERT (handles document upsert)
- timeout=300s on ANN queries and bulk INSERT for large banks
- Demote [ANN] debug logs to logger.debug()
- Fix docstring typos (agent_id → bank_id)
- Fix content_index remapping in producer-consumer batches
- Fix delta retain passing contents vs delta_contents
50MB benchmark (mock LLM): 9.2 min (was 23 min) — 2.5x faster.
BEAM 10m benchmark: zero deadlocks, zero DB errors.
* refactor(retain): remove legacy fallback code paths
- Remove process_entities_batch (legacy single-connection entity processing)
- Remove extract_entities_batch_optimized (only caller was the above)
- Remove fallback entity processing inside Phase 2 transaction
- Remove legacy ANN inline fallback in create_semantic_links_batch
- Remove fallback entity_links direct-insert path in Phase 3
- Make resolved_entity_ids/entity_to_unit/unit_to_entity_ids required params
* refactor(retain): replace tuple returns with dataclasses, remove dead code
- Add EntityResolutionResult and Phase1Result dataclasses in types.py
- Replace 4-tuple return from _pre_resolve_phase1 with Phase1Result
- Remove dead `entity_links = []` variables in retain_batch and _try_delta_retain
- Remove unused `confidence_score` parameter from orchestrator.retain_batch
and _retain_batch_async_internal (was accepted but never used)
* fix(entity-resolver): remove LIKE full-scan fallbacks, use index-only trigram matching
The entity resolution query had LIKE '%...' substring conditions that bypassed
the GIN trigram index, causing full sequential scans of the entities table.
On banks with 10k+ entities, this caused TimeoutErrors (observed in BEAM 10m).
Changes:
- Remove LIKE fallbacks, use trigram % operator only (GIN index-based)
- Lower similarity threshold from 0.3 to 0.15 to catch substring relationships
- Use LOWER() on both sides for case-insensitive matching
- Migration: recreate GIN trigram index on LOWER(canonical_name)
* fix: remove schema prefix from index names in trigram migration
* fix(delta-retain): use same chunk_size as streaming path (3000 vs 120000)
_chunk_contents_for_delta defaulted to chunk_size=120000 while the streaming
path used 3000. On retry, delta re-chunked the document with different
boundaries, found 0 matching chunks, and fell through to full re-extraction.
This wasted all LLM calls on already-committed chunks.
Fix: use the same default (3000) so chunk hashes match on recovery.
* fix(retain): persist generated document_id in operation metadata for retry recovery
When no document_id is provided, retain generates a UUID. On retry, a new UUID
was generated, making delta retain and streaming chunk-hash recovery unable to
find previously committed chunks. All LLM extraction was wasted on retry.
Fix: resolve document_id early in retain_batch (before delta), persist it to
operation result_metadata, and recover it on retry. Both delta and streaming
paths now see the same document_id across attempts.
* refactor(retain): unify into single streaming pipeline, remove non-streaming path
All retains now go through the producer-consumer streaming pipeline,
regardless of document size. Small documents are processed as a single batch.
This eliminates the maintenance burden of two separate code paths.
Also fix document upsert: compare content hash to distinguish recovery
(same content, partially committed) from update (different content, needs
cascade-delete). Previously, existing chunks always triggered recovery mode.
* refactor(retain): remove dead code, replace raw dicts with Phase3Context dataclass
- Remove dead _handle_zero_facts_documents (no callers after path unification)
- Remove unused imports: defaultdict, EntityLink
- Replace raw dict phase3_context with typed Phase3Context dataclass
- Update _build_and_insert_entity_links_phase3 to use typed parameter
3335 lines
132 KiB
Python
3335 lines
132 KiB
Python
"""
|
||
Test retain function and chunk storage.
|
||
"""
|
||
import logging
|
||
from datetime import datetime, timedelta, timezone
|
||
|
||
import pytest
|
||
|
||
from hindsight_api import RequestContext
|
||
from hindsight_api.engine.memory_engine import Budget
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_retain_with_chunks(memory, request_context):
|
||
"""
|
||
Test that retain function:
|
||
1. Stores facts with associated chunks
|
||
2. Recall returns chunk_id for each fact
|
||
"""
|
||
bank_id = f"test_chunks_{datetime.now(timezone.utc).timestamp()}"
|
||
document_id = "test_doc_123"
|
||
|
||
try:
|
||
# Store content that will be chunked (long enough to create multiple facts)
|
||
long_content = """
|
||
Alice is a senior software engineer at TechCorp. She has been working there for 5 years.
|
||
Alice specializes in distributed systems and has led the development of the company's
|
||
microservices architecture. She is known for writing clean, well-documented code.
|
||
|
||
Bob joined the team last month as a junior developer. He is learning React and Node.js.
|
||
Bob is enthusiastic and asks great questions during code reviews. He recently completed
|
||
his first feature, which was a user authentication flow.
|
||
|
||
The team uses Kubernetes for container orchestration and deploys to AWS. They follow
|
||
agile methodologies with two-week sprints. Code reviews are mandatory before merging.
|
||
"""
|
||
|
||
# Retain with document_id to enable chunk storage
|
||
unit_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=long_content,
|
||
context="team overview",
|
||
event_date=datetime(2024, 1, 15, tzinfo=timezone.utc),
|
||
document_id=document_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
print(f"\n=== Retained {len(unit_ids)} facts ===")
|
||
assert len(unit_ids) > 0, "Should have extracted and stored facts"
|
||
|
||
# Test 1: Recall with chunks enabled
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="Tell me about Alice",
|
||
budget=Budget.LOW,
|
||
max_tokens=500,
|
||
fact_type=["world"], # Search for world facts
|
||
include_chunks=True, # Enable chunks
|
||
max_chunk_tokens=8192,
|
||
request_context=request_context,
|
||
)
|
||
|
||
print("\n=== Recall Results (with chunks) ===")
|
||
print(f"Found {len(result.results)} results")
|
||
|
||
assert len(result.results) > 0, "Should find facts about Alice"
|
||
|
||
# Verify that chunks are returned
|
||
assert result.chunks is not None, "Chunks should be included in the response"
|
||
assert len(result.chunks) > 0, "Should have at least one chunk"
|
||
|
||
print(f"Number of chunks returned: {len(result.chunks)}")
|
||
|
||
# Verify chunk structure
|
||
for chunk_id, chunk_info in result.chunks.items():
|
||
print(f"\nChunk {chunk_id}:")
|
||
print(f" - chunk_index: {chunk_info.chunk_index}")
|
||
print(f" - chunk_text length: {len(chunk_info.chunk_text)} chars")
|
||
print(f" - truncated: {chunk_info.truncated}")
|
||
print(f" - text preview: {chunk_info.chunk_text[:100]}...")
|
||
|
||
# Verify chunk structure
|
||
assert isinstance(chunk_info.chunk_index, int), "Chunk index should be an integer"
|
||
assert chunk_info.chunk_index >= 0, "Chunk index should be non-negative"
|
||
assert len(chunk_info.chunk_text) > 0, "Chunk text should not be empty"
|
||
assert isinstance(chunk_info.truncated, bool), "Truncated should be boolean"
|
||
|
||
print("\n=== Test passed: Chunks are stored and retrieved correctly ===")
|
||
|
||
finally:
|
||
# Cleanup - delete the test bank
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
print(f"\n=== Cleaned up bank: {bank_id} ===")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_chunks_and_entities_follow_fact_order(memory, request_context):
|
||
"""
|
||
Test that chunks and entities in recall results follow the same order as facts.
|
||
This is critical because token limits may truncate later items.
|
||
|
||
The most relevant fact's chunk/entity should always be first in the returned data.
|
||
"""
|
||
bank_id = f"test_ordering_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store multiple distinct pieces of content as separate documents
|
||
# This ensures different chunks that we can identify
|
||
contents = [
|
||
{
|
||
"content": "Alice works at Google as a software engineer. She loves Python and has 10 years of experience.",
|
||
"document_id": "doc_alice",
|
||
"context": "Alice's profile"
|
||
},
|
||
{
|
||
"content": "Bob works at Meta as a data scientist. He specializes in machine learning and has published papers.",
|
||
"document_id": "doc_bob",
|
||
"context": "Bob's profile"
|
||
},
|
||
{
|
||
"content": "Charlie works at Amazon as a product manager. He leads a team of 15 people and ships features weekly.",
|
||
"document_id": "doc_charlie",
|
||
"context": "Charlie's profile"
|
||
},
|
||
]
|
||
|
||
# Store each content piece
|
||
for item in contents:
|
||
await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=item["content"],
|
||
context=item["context"],
|
||
event_date=datetime(2024, 1, 15, tzinfo=timezone.utc),
|
||
document_id=item["document_id"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
print("\n=== Stored 3 separate documents ===")
|
||
|
||
# Recall with a query that matches all three, but Alice most closely
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="Tell me about Alice's work at Google",
|
||
budget=Budget.MID,
|
||
max_tokens=1000,
|
||
fact_type=["world"],
|
||
include_chunks=True,
|
||
max_chunk_tokens=8192,
|
||
request_context=request_context,
|
||
)
|
||
|
||
print("\n=== Recall Results ===")
|
||
print(f"Found {len(result.results)} facts")
|
||
|
||
# Extract the order of entities mentioned in facts
|
||
fact_chunk_ids = []
|
||
fact_entities = []
|
||
|
||
for i, fact in enumerate(result.results):
|
||
print(f"\nFact {i}: {fact.text[:80]}...")
|
||
print(f" chunk_id: {fact.chunk_id}")
|
||
|
||
# Track chunk_id order
|
||
if fact.chunk_id:
|
||
fact_chunk_ids.append(fact.chunk_id)
|
||
|
||
# Track entities mentioned in this fact
|
||
if fact.entities:
|
||
for entity in fact.entities:
|
||
if entity not in fact_entities:
|
||
fact_entities.append(entity)
|
||
|
||
print(f"\n=== Fact chunk_ids in order: {fact_chunk_ids} ===")
|
||
print(f"=== Fact entities in order: {fact_entities} ===")
|
||
|
||
# Test 1: Verify chunks follow fact order
|
||
if result.chunks:
|
||
chunks_order = list(result.chunks.keys())
|
||
print(f"\n=== Chunks dict order: {chunks_order} ===")
|
||
|
||
# The chunks dict should contain chunks in the order they appear in facts
|
||
# (may be fewer chunks than facts due to deduplication)
|
||
chunk_positions = []
|
||
for chunk_id in chunks_order:
|
||
if chunk_id in fact_chunk_ids:
|
||
chunk_positions.append(fact_chunk_ids.index(chunk_id))
|
||
|
||
print(f"=== Chunk positions in fact order: {chunk_positions} ===")
|
||
|
||
# Verify chunks are in increasing order (following fact order)
|
||
assert chunk_positions == sorted(chunk_positions), \
|
||
f"Chunks should follow fact order! Got positions {chunk_positions} but expected {sorted(chunk_positions)}"
|
||
|
||
print("✓ Chunks follow fact order correctly")
|
||
|
||
# Test 2: Verify entities follow fact order
|
||
if result.entities:
|
||
entities_order = list(result.entities.keys())
|
||
print(f"\n=== Entities dict order: {entities_order} ===")
|
||
|
||
# The entities dict should contain entities in the order they first appear in facts
|
||
entity_positions = []
|
||
for entity_name in entities_order:
|
||
if entity_name in fact_entities:
|
||
entity_positions.append(fact_entities.index(entity_name))
|
||
|
||
print(f"=== Entity positions in fact order: {entity_positions} ===")
|
||
|
||
# Verify entities are in increasing order (following fact order)
|
||
assert entity_positions == sorted(entity_positions), \
|
||
f"Entities should follow fact order! Got positions {entity_positions} but expected {sorted(entity_positions)}"
|
||
|
||
print("✓ Entities follow fact order correctly")
|
||
|
||
print("\n=== Test passed: Chunks and entities follow fact relevance order ===")
|
||
|
||
finally:
|
||
# Cleanup
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
print(f"\n=== Cleaned up bank: {bank_id} ===")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_event_date_storage(memory, request_context):
|
||
"""
|
||
Test that event_date is correctly stored as occurred_start.
|
||
Verifies that we can track when events actually happened vs when they were stored.
|
||
"""
|
||
bank_id = f"test_temporal_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Event that occurred in the past
|
||
past_event_date = datetime(2023, 6, 15, 14, 30, tzinfo=timezone.utc)
|
||
|
||
# Store a fact about a past event
|
||
unit_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice completed the Q2 product launch on June 15th, 2023.",
|
||
context="project history",
|
||
event_date=past_event_date,
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids) > 0, "Should have created at least one memory unit"
|
||
|
||
# Recall the fact
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="When did Alice complete the product launch?",
|
||
budget=Budget.LOW,
|
||
max_tokens=500,
|
||
fact_type=["world"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should recall the stored fact"
|
||
|
||
# Verify the occurred_start matches our event_date
|
||
fact = result.results[0]
|
||
assert fact.occurred_start is not None, "occurred_start should be set"
|
||
|
||
# Parse the occurred_start (it comes back as ISO string)
|
||
if isinstance(fact.occurred_start, str):
|
||
occurred_dt = datetime.fromisoformat(fact.occurred_start.replace('Z', '+00:00'))
|
||
else:
|
||
occurred_dt = fact.occurred_start
|
||
|
||
# Verify it matches our past event date (allowing for small time differences in extraction)
|
||
assert occurred_dt.year == past_event_date.year, f"Year should match: {occurred_dt.year} vs {past_event_date.year}"
|
||
assert occurred_dt.month == past_event_date.month, f"Month should match: {occurred_dt.month} vs {past_event_date.month}"
|
||
|
||
print(f"\n✓ Event date correctly stored: {occurred_dt}")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.xfail(reason="LLM date extraction from content is non-deterministic", strict=False)
|
||
async def test_temporal_ordering(memory, request_context):
|
||
"""
|
||
Test that facts can be stored and retrieved with correct temporal ordering.
|
||
Stores facts with different event_dates and verifies temporal relationships.
|
||
"""
|
||
bank_id = f"test_temporal_order_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store events in non-chronological order with different dates
|
||
events = [
|
||
{
|
||
"content": "Alice joined the team in January 2023.",
|
||
"event_date": datetime(2023, 1, 10, tzinfo=timezone.utc),
|
||
"context": "team history"
|
||
},
|
||
{
|
||
"content": "Alice got promoted to senior engineer in June 2023.",
|
||
"event_date": datetime(2023, 6, 15, tzinfo=timezone.utc),
|
||
"context": "team history"
|
||
},
|
||
{
|
||
"content": "Alice started as an intern in July 2022.",
|
||
"event_date": datetime(2022, 7, 1, tzinfo=timezone.utc),
|
||
"context": "team history"
|
||
},
|
||
]
|
||
|
||
# Store all events
|
||
for event in events:
|
||
await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=event["content"],
|
||
context=event["context"],
|
||
event_date=event["event_date"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
print("\n=== Stored 3 events with different temporal dates ===")
|
||
|
||
# Recall facts about Alice
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="Tell me about Alice's career progression",
|
||
budget=Budget.MID,
|
||
max_tokens=1000,
|
||
fact_type=["world"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) >= 2, f"Should recall at least 2 events, got {len(result.results)}"
|
||
|
||
# Collect occurred dates
|
||
occurred_dates = []
|
||
for fact in result.results:
|
||
if fact.occurred_start:
|
||
if isinstance(fact.occurred_start, str):
|
||
dt = datetime.fromisoformat(fact.occurred_start.replace('Z', '+00:00'))
|
||
else:
|
||
dt = fact.occurred_start
|
||
occurred_dates.append((dt, fact.text[:50]))
|
||
print(f" - {dt.date()}: {fact.text[:60]}...")
|
||
|
||
# Verify we have temporal data for most facts (LLM may occasionally miss one)
|
||
assert len(occurred_dates) >= 2, "At least 2 facts should have temporal data"
|
||
|
||
# The dates should span the expected range (2022-2023)
|
||
min_date = min(dt for dt, _ in occurred_dates)
|
||
max_date = max(dt for dt, _ in occurred_dates)
|
||
|
||
assert min_date.year == 2022, f"Earliest event should be in 2022, got {min_date.year}"
|
||
assert max_date.year == 2023, f"Latest event should be in 2023, got {max_date.year}"
|
||
|
||
print(f"\n✓ Temporal ordering preserved: {min_date.date()} to {max_date.date()}")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_mentioned_at_vs_occurred(memory, request_context):
|
||
"""
|
||
Test distinction between when fact occurred vs when it was mentioned.
|
||
|
||
Scenario: Ingesting a historical conversation from 2020
|
||
- event_date: When the conversation happened (2020-03-15)
|
||
- mentioned_at: When the conversation happened (same as event_date = 2020-03-15)
|
||
- occurred_start/end: When the event in the conversation happened (extracted by LLM, or falls back to mentioned_at)
|
||
"""
|
||
bank_id = f"test_mentioned_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Ingesting a conversation that happened in the past
|
||
conversation_date = datetime(2020, 3, 15, tzinfo=timezone.utc)
|
||
|
||
# Store a fact from a historical conversation
|
||
unit_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice graduated from MIT in March 2020.",
|
||
context="education history",
|
||
event_date=conversation_date, # When this conversation happened
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids) > 0, "Should create memory unit"
|
||
|
||
# Recall and check temporal fields
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="Where did Alice go to school?",
|
||
budget=Budget.LOW,
|
||
max_tokens=500,
|
||
fact_type=["world"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should recall the fact"
|
||
fact = result.results[0]
|
||
|
||
# Parse occurred_start
|
||
if fact.occurred_start:
|
||
if isinstance(fact.occurred_start, str):
|
||
occurred_dt = datetime.fromisoformat(fact.occurred_start.replace('Z', '+00:00'))
|
||
else:
|
||
occurred_dt = fact.occurred_start
|
||
|
||
# Should be close to the conversation date (falls back to mentioned_at if LLM doesn't extract)
|
||
assert occurred_dt.year == 2020, f"occurred_start should be 2020, got {occurred_dt.year}"
|
||
print(f"✓ occurred_start (when event happened): {occurred_dt}")
|
||
|
||
# Parse mentioned_at
|
||
if fact.mentioned_at:
|
||
if isinstance(fact.mentioned_at, str):
|
||
mentioned_dt = datetime.fromisoformat(fact.mentioned_at.replace('Z', '+00:00'))
|
||
else:
|
||
mentioned_dt = fact.mentioned_at
|
||
|
||
# mentioned_at should match the conversation date (event_date)
|
||
time_diff = abs((conversation_date - mentioned_dt).total_seconds())
|
||
assert time_diff < 60, f"mentioned_at should match event_date (2020-03-15), but diff is {time_diff}s"
|
||
print(f"✓ mentioned_at (when conversation happened): {mentioned_dt}")
|
||
|
||
# Verify it's the historical date, not today
|
||
assert mentioned_dt.year == 2020, f"mentioned_at should be 2020, got {mentioned_dt.year}"
|
||
|
||
print("✓ Test passed: Historical conversation correctly ingested with event_date=2020")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_occurred_dates_not_defaulted(memory, request_context):
|
||
"""
|
||
Test that occurred_start and occurred_end are NOT defaulted to mentioned_at.
|
||
|
||
This is a regression test for a bug where occurred dates were incorrectly
|
||
defaulting to mentioned_at when the LLM didn't provide them.
|
||
|
||
Scenario: Store a fact where occurred dates are not applicable (current observation)
|
||
- mentioned_at should be set (to event_date or now())
|
||
- occurred_start and occurred_end should be None (not defaulted to mentioned_at)
|
||
"""
|
||
bank_id = f"test_occurred_not_defaulted_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store a current observation where occurred dates don't make sense
|
||
# Use present tense to avoid LLM extracting past dates
|
||
# Content needs to be substantial enough to not be filtered as trivial
|
||
event_date = datetime(2024, 2, 10, 15, 30, tzinfo=timezone.utc)
|
||
|
||
unit_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice is a software engineer who specializes in Python and machine learning. She prefers dark roast coffee and works remotely from Seattle.",
|
||
context="current observations about Alice",
|
||
event_date=event_date,
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids) > 0, "Should create memory unit"
|
||
|
||
# Recall and check that occurred dates are None
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="Tell me about Alice",
|
||
budget=Budget.LOW,
|
||
max_tokens=500,
|
||
fact_type=["world", "experience"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should recall the fact"
|
||
fact = result.results[0]
|
||
|
||
# mentioned_at should be set
|
||
assert fact.mentioned_at is not None, "mentioned_at should be set"
|
||
|
||
# Parse mentioned_at
|
||
if isinstance(fact.mentioned_at, str):
|
||
mentioned_dt = datetime.fromisoformat(fact.mentioned_at.replace('Z', '+00:00'))
|
||
else:
|
||
mentioned_dt = fact.mentioned_at
|
||
|
||
# Verify it matches event_date
|
||
time_diff = abs((event_date - mentioned_dt).total_seconds())
|
||
assert time_diff < 60, f"mentioned_at should match event_date, but diff is {time_diff}s"
|
||
|
||
# CRITICAL: occurred_start and occurred_end should be None
|
||
# They should NOT default to mentioned_at
|
||
if fact.occurred_start is not None:
|
||
# If occurred_start is set, it means the LLM extracted it
|
||
# In this case, log it but don't fail (LLM behavior can vary)
|
||
print(f"⚠ LLM extracted occurred_start: {fact.occurred_start}")
|
||
print(" This test expects None for present-tense observations")
|
||
else:
|
||
print("✓ occurred_start is correctly None (not defaulted to mentioned_at)")
|
||
|
||
if fact.occurred_end is not None:
|
||
print(f"⚠ LLM extracted occurred_end: {fact.occurred_end}")
|
||
print(" This test expects None for present-tense observations")
|
||
else:
|
||
print("✓ occurred_end is correctly None (not defaulted to mentioned_at)")
|
||
|
||
# At least verify they're not equal to mentioned_at if they are set
|
||
if fact.occurred_start is not None:
|
||
if isinstance(fact.occurred_start, str):
|
||
occurred_start_dt = datetime.fromisoformat(fact.occurred_start.replace('Z', '+00:00'))
|
||
else:
|
||
occurred_start_dt = fact.occurred_start
|
||
|
||
# If they're equal, it suggests the old defaulting bug
|
||
if occurred_start_dt == mentioned_dt:
|
||
raise AssertionError(
|
||
f"occurred_start should NOT be defaulted to mentioned_at! "
|
||
f"occurred_start={occurred_start_dt}, mentioned_at={mentioned_dt}"
|
||
)
|
||
|
||
print("✓ Test passed: occurred dates are not incorrectly defaulted to mentioned_at")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_mentioned_at_from_context_string(memory, request_context):
|
||
"""
|
||
Test that mentioned_at is extracted from context string by LLM.
|
||
|
||
Scenario: User provides date in context like "happened on 2023-05-10 14:30:00 UTC"
|
||
- LLM should extract mentioned_at from this context
|
||
- If LLM fails to extract, should fall back to event_date (which defaults to now())
|
||
- mentioned_at should NEVER be None
|
||
"""
|
||
bank_id = f"test_context_date_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Test case 1: Date in context string (like longmemeval benchmark)
|
||
session_date = datetime(2023, 5, 10, 14, 30, 0, tzinfo=timezone.utc)
|
||
|
||
unit_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice mentioned she loves hiking in the mountains.",
|
||
context=f"Session ABC123 - you are the assistant in this conversation - happened on {session_date.strftime('%Y-%m-%d %H:%M:%S')} UTC.",
|
||
event_date=None, # Not providing event_date - should default to now() if LLM doesn't extract
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids) > 0, "Should create memory unit"
|
||
|
||
# Recall and verify mentioned_at is set
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="What does Alice like?",
|
||
budget=Budget.LOW,
|
||
max_tokens=500,
|
||
fact_type=["world"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should recall the fact"
|
||
fact = result.results[0]
|
||
|
||
# mentioned_at must ALWAYS be set
|
||
assert fact.mentioned_at is not None, "mentioned_at should NEVER be None"
|
||
|
||
# Parse mentioned_at
|
||
if isinstance(fact.mentioned_at, str):
|
||
mentioned_dt = datetime.fromisoformat(fact.mentioned_at.replace('Z', '+00:00'))
|
||
else:
|
||
mentioned_dt = fact.mentioned_at
|
||
|
||
# Check if LLM extracted the date from context (ideal case)
|
||
# Or if it fell back to now() (acceptable fallback)
|
||
time_diff_from_context = abs((session_date - mentioned_dt).total_seconds())
|
||
time_diff_from_now = abs((datetime.now(timezone.utc) - mentioned_dt).total_seconds())
|
||
|
||
# Should either match the context date OR be recent (now)
|
||
is_from_context = time_diff_from_context < 60
|
||
is_from_now = time_diff_from_now < 60
|
||
|
||
assert is_from_context or is_from_now, \
|
||
f"mentioned_at should be either from context ({session_date}) or now(), but got {mentioned_dt}"
|
||
|
||
if is_from_context:
|
||
print(f"✓ LLM successfully extracted mentioned_at from context: {mentioned_dt}")
|
||
assert mentioned_dt.year == 2023
|
||
else:
|
||
print(f"⚠ LLM did not extract date from context, fell back to now(): {mentioned_dt}")
|
||
|
||
print("✓ mentioned_at is always set (never None)")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
# ============================================================
|
||
# No Timestamp Tests
|
||
# ============================================================
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_retain_no_timestamp(memory, request_context):
|
||
"""
|
||
Test retaining content with explicit "no timestamp" sentinel.
|
||
|
||
When event_date=None is passed explicitly in the dict (i.e. caller opted into
|
||
no timestamp), mentioned_at should be NULL in the DB rather than defaulting to now().
|
||
"""
|
||
bank_id = f"test_no_timestamp_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Use retain_batch_async with explicit event_date=None key to signal "no timestamp"
|
||
unit_ids_list = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[
|
||
{
|
||
"content": "The capital of France is Paris. The Eiffel Tower is located in Paris.",
|
||
"context": "general knowledge",
|
||
"event_date": None, # Explicit sentinel: no timestamp
|
||
}
|
||
],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids_list) > 0, "Should create at least one batch result"
|
||
unit_ids = unit_ids_list[0]
|
||
assert len(unit_ids) > 0, "Should have extracted and stored facts"
|
||
|
||
# Recall the facts
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="Where is the Eiffel Tower?",
|
||
budget=Budget.LOW,
|
||
max_tokens=500,
|
||
fact_type=["world"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should recall the stored fact"
|
||
|
||
# All temporal fields should be None for temporally agnostic content
|
||
for fact in result.results:
|
||
assert fact.mentioned_at is None, (
|
||
f"mentioned_at should be None for no-timestamp content, got {fact.mentioned_at}"
|
||
)
|
||
|
||
print(f"\n✓ Test passed: mentioned_at is None for {len(result.results)} fact(s)")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_retain_omit_timestamp_defaults_to_now(memory, request_context):
|
||
"""
|
||
Backward-compatibility regression test: omitting event_date still stores a real datetime.
|
||
|
||
When event_date is absent from the content dict (key not present), the orchestrator
|
||
should default to utcnow() — preserving existing behavior.
|
||
"""
|
||
bank_id = f"test_default_timestamp_{datetime.now(timezone.utc).timestamp()}"
|
||
before = datetime.now(timezone.utc)
|
||
|
||
try:
|
||
# Omit event_date entirely — should default to now()
|
||
unit_ids_list = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[
|
||
{
|
||
"content": "Alice is a software engineer who loves Python.",
|
||
"context": "profile",
|
||
# event_date intentionally omitted
|
||
}
|
||
],
|
||
request_context=request_context,
|
||
)
|
||
|
||
after = datetime.now(timezone.utc)
|
||
|
||
assert len(unit_ids_list) > 0
|
||
unit_ids = unit_ids_list[0]
|
||
assert len(unit_ids) > 0, "Should have extracted and stored facts"
|
||
|
||
# Recall and verify mentioned_at is a real datetime close to now
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="Who is Alice?",
|
||
budget=Budget.LOW,
|
||
max_tokens=500,
|
||
fact_type=["world"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should recall the fact"
|
||
fact = result.results[0]
|
||
|
||
assert fact.mentioned_at is not None, "mentioned_at should be set when event_date is omitted"
|
||
|
||
if isinstance(fact.mentioned_at, str):
|
||
mentioned_dt = datetime.fromisoformat(fact.mentioned_at.replace("Z", "+00:00"))
|
||
else:
|
||
mentioned_dt = fact.mentioned_at
|
||
|
||
# Should be within 60s of when we ran the test
|
||
assert before <= mentioned_dt <= after + timedelta(seconds=60), (
|
||
f"mentioned_at {mentioned_dt} should be close to now ({before} – {after})"
|
||
)
|
||
|
||
print(f"\n✓ Test passed: mentioned_at={mentioned_dt} is a real datetime (backward compat)")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
# ============================================================
|
||
# Context Tracking Tests
|
||
# ============================================================
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_context_preservation(memory, request_context):
|
||
"""
|
||
Test that context is preserved and retrievable.
|
||
Context helps understand why/how memory was formed.
|
||
"""
|
||
bank_id = f"test_context_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store content with specific context
|
||
specific_context = "team meeting notes from Q4 planning session"
|
||
|
||
unit_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="The team decided to prioritize mobile development for next quarter.",
|
||
context=specific_context,
|
||
event_date=datetime(2024, 1, 15, tzinfo=timezone.utc),
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids) > 0, "Should create at least one memory unit"
|
||
|
||
# Recall and verify context is returned
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="What did the team decide?",
|
||
budget=Budget.LOW,
|
||
max_tokens=500,
|
||
fact_type=["world"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should recall the stored fact"
|
||
|
||
# Verify context is preserved (context is stored in the database)
|
||
# Note: context might not be returned in the API response by default
|
||
# but it should be stored in the database
|
||
print(f"✓ Successfully stored fact with context: '{specific_context}'")
|
||
print(f" Retrieved {len(result.results)} facts")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_context_with_batch(memory, request_context):
|
||
"""
|
||
Test that each item in a batch can have different contexts.
|
||
|
||
Note: LLM fact extraction is non-deterministic. Simple sentences may
|
||
not always produce exactly 1 fact each. We verify the batch was
|
||
processed and at least some facts were extracted.
|
||
"""
|
||
bank_id = f"test_batch_context_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store batch with different contexts
|
||
unit_ids = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[
|
||
{
|
||
"content": "Alice completed the authentication module.",
|
||
"context": "sprint 1 standup",
|
||
"event_date": datetime(2024, 1, 10, tzinfo=timezone.utc)
|
||
},
|
||
{
|
||
"content": "Bob started working on the database schema.",
|
||
"context": "sprint 1 planning",
|
||
"event_date": datetime(2024, 1, 11, tzinfo=timezone.utc)
|
||
},
|
||
{
|
||
"content": "Charlie fixed critical bugs in the payment flow.",
|
||
"context": "incident response",
|
||
"event_date": datetime(2024, 1, 12, tzinfo=timezone.utc)
|
||
}
|
||
],
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Should have created facts from at least some items
|
||
# LLM extraction is non-deterministic, so we allow some flexibility
|
||
total_units = sum(len(ids) for ids in unit_ids)
|
||
assert total_units >= 2, f"Should create at least 2 units from 3 batch items, got {total_units}"
|
||
|
||
print(f"✓ Stored {len(unit_ids)} batch items with different contexts")
|
||
print(f" Created {total_units} total memory units")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
# ============================================================
|
||
# Metadata Storage Tests
|
||
# ============================================================
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_metadata_storage_and_retrieval(memory, request_context):
|
||
"""
|
||
Test that user-defined metadata passed during retain is returned on recall.
|
||
Metadata allows arbitrary key-value data to be stored with facts.
|
||
"""
|
||
bank_id = f"test_metadata_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
custom_metadata = {
|
||
"source": "slack",
|
||
"channel": "engineering",
|
||
"importance": "high",
|
||
}
|
||
|
||
# Use retain_batch_async which supports the metadata parameter
|
||
unit_ids_list = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[
|
||
{
|
||
"content": "The product launch is scheduled for March 1st.",
|
||
"context": "planning meeting",
|
||
"event_date": datetime(2024, 1, 15, tzinfo=timezone.utc),
|
||
"metadata": custom_metadata,
|
||
}
|
||
],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids_list) > 0, "Should create memory units"
|
||
assert len(unit_ids_list[0]) > 0, "Should have at least one unit ID"
|
||
|
||
# Recall and verify metadata is returned
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="When is the product launch?",
|
||
budget=Budget.LOW,
|
||
max_tokens=500,
|
||
fact_type=["world"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should recall stored facts"
|
||
|
||
# Verify metadata is present on recalled facts
|
||
fact = result.results[0]
|
||
assert fact.metadata is not None, "Metadata should not be null on recall"
|
||
assert fact.metadata.get("source") == "slack"
|
||
assert fact.metadata.get("channel") == "engineering"
|
||
assert fact.metadata.get("importance") == "high"
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
# ============================================================
|
||
# Batch Processing Edge Cases
|
||
# ============================================================
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_empty_batch(memory, request_context):
|
||
"""
|
||
Test that empty batch is handled gracefully without errors.
|
||
"""
|
||
bank_id = f"test_empty_batch_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Attempt to store empty batch
|
||
unit_ids = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[],
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Should return empty list or handle gracefully
|
||
assert isinstance(unit_ids, list), "Should return a list"
|
||
assert len(unit_ids) == 0, "Empty batch should create no units"
|
||
|
||
print("✓ Empty batch handled gracefully")
|
||
|
||
finally:
|
||
# Clean up (though nothing should be stored)
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_single_item_batch(memory, request_context):
|
||
"""
|
||
Test that batch with one item works correctly.
|
||
"""
|
||
bank_id = f"test_single_batch_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store batch with single item
|
||
unit_ids = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[
|
||
{
|
||
"content": "Alice shipped the new feature to production.",
|
||
"context": "deployment log",
|
||
"event_date": datetime(2024, 1, 15, tzinfo=timezone.utc)
|
||
}
|
||
],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids) == 1, "Should return one list of unit IDs"
|
||
assert len(unit_ids[0]) > 0, "Should create at least one memory unit"
|
||
|
||
print(f"✓ Single-item batch created {len(unit_ids[0])} units")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_mixed_content_batch(memory, request_context):
|
||
"""
|
||
Test batch with varying content sizes (short and long).
|
||
"""
|
||
bank_id = f"test_mixed_batch_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Mix short and long content
|
||
short_content = "Alice joined the team."
|
||
long_content = """
|
||
Bob has been working on the authentication system for the past three months.
|
||
He implemented OAuth 2.0 integration, set up JWT token management, and built
|
||
a comprehensive role-based access control system. The system supports multiple
|
||
identity providers including Google, GitHub, and Microsoft. Bob also wrote
|
||
extensive documentation and unit tests covering over 90% of the codebase.
|
||
The team recognized his work with an excellence award at the quarterly meeting.
|
||
"""
|
||
|
||
unit_ids = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[
|
||
{"content": short_content, "context": "onboarding"},
|
||
{"content": long_content, "context": "performance review"},
|
||
{"content": "Charlie is on vacation this week.", "context": "team status"}
|
||
],
|
||
request_context=request_context,
|
||
)
|
||
|
||
# All items should be processed
|
||
assert len(unit_ids) == 3, "Should process all 3 items"
|
||
|
||
# Long content should create more facts
|
||
short_units = len(unit_ids[0])
|
||
long_units = len(unit_ids[1])
|
||
|
||
print("✓ Mixed batch processed successfully")
|
||
print(f" Short content: {short_units} units")
|
||
print(f" Long content: {long_units} units")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_batch_with_missing_optional_fields(memory, request_context):
|
||
"""
|
||
Test that batch handles items with missing optional fields.
|
||
"""
|
||
bank_id = f"test_optional_fields_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Some items have all fields, some have minimal fields
|
||
unit_ids = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[
|
||
{
|
||
"content": "Alice finished the project.",
|
||
"context": "complete record",
|
||
"event_date": datetime(2024, 1, 15, tzinfo=timezone.utc)
|
||
},
|
||
{
|
||
"content": "Bob started a new task.",
|
||
# No context or event_date
|
||
},
|
||
{
|
||
"content": "Charlie reviewed code.",
|
||
"context": "code review",
|
||
# No event_date
|
||
}
|
||
],
|
||
request_context=request_context,
|
||
)
|
||
|
||
# All items should be processed successfully
|
||
assert len(unit_ids) == 3, "Should process all items even with missing optional fields"
|
||
|
||
total_units = sum(len(ids) for ids in unit_ids)
|
||
print(f"✓ Batch with mixed optional fields created {total_units} total units")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
# ============================================================
|
||
# Multi-Document Batch Tests
|
||
# ============================================================
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_single_batch_multiple_documents(memory, request_context):
|
||
"""
|
||
Test storing multiple distinct documents in a single batch call.
|
||
Each should be tracked separately.
|
||
"""
|
||
bank_id = f"test_multi_docs_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store single batch where each item could be a different document
|
||
# (In practice, document_id is a batch-level parameter, so we test
|
||
# that multiple retain_async calls work correctly)
|
||
|
||
doc1_units = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice's resume: 10 years Python experience, worked at Google.",
|
||
context="resume review",
|
||
document_id="resume_alice",
|
||
request_context=request_context,
|
||
)
|
||
|
||
doc2_units = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Bob's resume: 5 years JavaScript experience, worked at Meta.",
|
||
context="resume review",
|
||
document_id="resume_bob",
|
||
request_context=request_context,
|
||
)
|
||
|
||
doc3_units = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Charlie's resume: 8 years Go experience, worked at Amazon.",
|
||
context="resume review",
|
||
document_id="resume_charlie",
|
||
request_context=request_context,
|
||
)
|
||
|
||
# All documents should be stored
|
||
assert len(doc1_units) > 0, "Should create units for doc1"
|
||
assert len(doc2_units) > 0, "Should create units for doc2"
|
||
assert len(doc3_units) > 0, "Should create units for doc3"
|
||
|
||
total_units = len(doc1_units) + len(doc2_units) + len(doc3_units)
|
||
print(f"✓ Stored 3 separate documents with {total_units} total units")
|
||
|
||
# Verify we can recall from any document
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="Who worked at Google?",
|
||
budget=Budget.MID,
|
||
max_tokens=1000,
|
||
fact_type=["world"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should find facts about Alice"
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_document_upsert_behavior(memory, request_context):
|
||
"""
|
||
Test that upserting a document replaces the old content.
|
||
"""
|
||
bank_id = f"test_upsert_{datetime.now(timezone.utc).timestamp()}"
|
||
document_id = "project_status"
|
||
|
||
try:
|
||
# Store initial version
|
||
v1_units = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Project is in planning phase. Alice is the lead.",
|
||
context="status update v1",
|
||
document_id=document_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(v1_units) > 0, "Should create units for v1"
|
||
|
||
# Update with new version (upsert)
|
||
v2_units = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Project is in development phase. Bob has joined as co-lead.",
|
||
context="status update v2",
|
||
document_id=document_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(v2_units) > 0, "Should create units for v2"
|
||
|
||
# Recall should return the updated information
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="What is the project status?",
|
||
budget=Budget.MID,
|
||
max_tokens=1000,
|
||
fact_type=["world"],
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should recall facts"
|
||
|
||
print(f"✓ Document upsert created v1: {len(v1_units)} units, v2: {len(v2_units)} units")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
# ============================================================
|
||
# Chunk Storage Advanced Tests
|
||
# ============================================================
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_chunk_fact_mapping(memory, request_context):
|
||
"""
|
||
Test that facts correctly reference their source chunks via chunk_id.
|
||
"""
|
||
bank_id = f"test_chunk_mapping_{datetime.now(timezone.utc).timestamp()}"
|
||
document_id = "technical_doc"
|
||
|
||
try:
|
||
# Store content that will be chunked
|
||
content = """
|
||
The authentication system uses JWT tokens for session management.
|
||
Tokens expire after 24 hours and must be refreshed using the refresh endpoint.
|
||
The system supports OAuth 2.0 integration with Google and GitHub.
|
||
|
||
The database layer uses PostgreSQL with connection pooling.
|
||
We maintain separate read and write connection pools for performance.
|
||
All queries use prepared statements to prevent SQL injection.
|
||
"""
|
||
|
||
unit_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=content,
|
||
context="technical documentation",
|
||
document_id=document_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids) > 0, "Should create memory units"
|
||
|
||
# Recall with chunks enabled
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="How does authentication work?",
|
||
budget=Budget.MID,
|
||
max_tokens=1000,
|
||
fact_type=["world"],
|
||
include_chunks=True,
|
||
max_chunk_tokens=8192,
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result.results) > 0, "Should recall facts"
|
||
|
||
# Verify facts have chunk_id references
|
||
facts_with_chunks = [f for f in result.results if f.chunk_id]
|
||
|
||
print(f"✓ Created {len(unit_ids)} units from chunked document")
|
||
print(f" {len(facts_with_chunks)}/{len(result.results)} facts have chunk_id references")
|
||
|
||
# If chunks are returned, verify they match the chunk_ids in facts
|
||
if result.chunks:
|
||
fact_chunk_ids = {f.chunk_id for f in facts_with_chunks}
|
||
returned_chunk_ids = set(result.chunks.keys())
|
||
|
||
# All chunk_ids in facts should have corresponding chunk data
|
||
assert fact_chunk_ids.issubset(returned_chunk_ids) or len(fact_chunk_ids) == 0, \
|
||
"Fact chunk_ids should have corresponding chunk data"
|
||
|
||
print(f" Returned {len(result.chunks)} chunks matching fact references")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_chunk_ordering_preservation(memory, request_context):
|
||
"""
|
||
Test that chunk_index reflects the correct order within a document.
|
||
"""
|
||
bank_id = f"test_chunk_order_{datetime.now(timezone.utc).timestamp()}"
|
||
document_id = "ordered_doc"
|
||
|
||
try:
|
||
# Store long content that will create multiple chunks with meaningful content
|
||
sections = []
|
||
sections.append("""
|
||
Alice is the team lead for the authentication project. She has 10 years of experience
|
||
with security systems and previously worked at Google on identity management.
|
||
She is responsible for architecture decisions and code review.
|
||
""")
|
||
sections.append("""
|
||
Bob is a backend engineer focusing on the API layer. He specializes in Python
|
||
and has built several microservices for the company. He joined the team in 2023.
|
||
""")
|
||
sections.append("""
|
||
Charlie is the DevOps engineer managing the deployment pipeline. He set up
|
||
our Kubernetes infrastructure and maintains the CI/CD system using GitHub Actions.
|
||
""")
|
||
sections.append("""
|
||
The project uses PostgreSQL as the main database with Redis for caching.
|
||
We deploy to AWS using Docker containers orchestrated by Kubernetes.
|
||
The team follows agile methodology with two-week sprints.
|
||
""")
|
||
sections.append("""
|
||
Security is a top priority. All API endpoints require JWT authentication.
|
||
We use OAuth 2.0 for third-party integrations and maintain strict access controls.
|
||
Regular security audits are conducted quarterly.
|
||
""")
|
||
|
||
content = "\n\n".join(sections)
|
||
|
||
unit_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=content,
|
||
context="multi-section document",
|
||
document_id=document_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids) > 0, "Should create units"
|
||
|
||
# Recall with chunks
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="Tell me about the sections",
|
||
budget=Budget.MID,
|
||
max_tokens=2000,
|
||
fact_type=["world"],
|
||
include_chunks=True,
|
||
max_chunk_tokens=8192,
|
||
request_context=request_context,
|
||
)
|
||
|
||
if result.chunks:
|
||
# Verify chunk_index values are sequential and start from 0
|
||
chunk_indices = [chunk.chunk_index for chunk in result.chunks.values()]
|
||
chunk_indices_sorted = sorted(chunk_indices)
|
||
|
||
print(f"✓ Document created {len(result.chunks)} chunks")
|
||
print(f" Chunk indices: {chunk_indices}")
|
||
|
||
# Indices should start from 0 and be sequential
|
||
if len(chunk_indices) > 0:
|
||
assert min(chunk_indices) == 0, "Chunk indices should start from 0"
|
||
assert chunk_indices_sorted == list(range(len(chunk_indices))), \
|
||
"Chunk indices should be sequential"
|
||
else:
|
||
print("✓ Content stored (may have created single chunk or no chunks returned)")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.timeout(180) # Allow up to 3 minutes for this test
|
||
async def test_chunks_truncation_behavior(memory, request_context):
|
||
"""
|
||
Test that when chunks exceed max_chunk_tokens, truncation is indicated.
|
||
|
||
Note: This test processes larger content and may take longer than typical tests.
|
||
"""
|
||
bank_id = f"test_chunk_truncation_{datetime.now(timezone.utc).timestamp()}"
|
||
document_id = "large_doc"
|
||
|
||
try:
|
||
# Create a moderately large document with meaningful content
|
||
# Reduced from * 5 to * 2 for faster execution while still testing truncation
|
||
large_content = """
|
||
The company's product roadmap for 2024 includes several major initiatives.
|
||
The engineering team is expanding to support these efforts.
|
||
|
||
Alice leads the authentication team, which is implementing OAuth 2.0 and JWT tokens.
|
||
The team has been working on this for six months and expects to launch in Q2.
|
||
Security is the top priority, with regular penetration testing scheduled.
|
||
|
||
Bob manages the API development team. They are building RESTful endpoints
|
||
for all major features including user management, billing, and analytics.
|
||
The team uses Python with FastAPI and deploys to AWS Lambda.
|
||
|
||
Charlie oversees the infrastructure team. They maintain Kubernetes clusters
|
||
across three AWS regions for high availability. The team also manages
|
||
the CI/CD pipeline using GitHub Actions and ArgoCD.
|
||
|
||
The data engineering team, led by Diana, processes millions of events daily.
|
||
They use Apache Kafka for streaming and Snowflake for analytics.
|
||
Real-time dashboards are built with Grafana and Prometheus.
|
||
|
||
The mobile team is building iOS and Android apps using React Native.
|
||
They are targeting a beta launch in Q3 with select customers.
|
||
Push notifications and offline support are key features.
|
||
|
||
The design team has created a new design system that will be rolled out
|
||
across all products. The system includes components for accessibility
|
||
and internationalization support for 12 languages.
|
||
|
||
Customer support is being enhanced with AI-powered chatbots.
|
||
The system can handle common queries and escalate complex issues to humans.
|
||
Average response time has improved by 40% since implementation.
|
||
|
||
The marketing team is planning a major campaign for the product launch.
|
||
They are working with influencers and planning webinars for enterprise customers.
|
||
Early feedback from beta users has been very positive.
|
||
|
||
Sales operations are being streamlined with new CRM integrations.
|
||
The team can now track leads more effectively and automate follow-ups.
|
||
Conversion rates have increased by 25% in the pilot program.
|
||
|
||
The finance team is implementing new budgeting tools for better forecasting.
|
||
They are also working on automated expense reporting and approval workflows.
|
||
This will save approximately 100 hours per month in manual work.
|
||
""" * 2 # Repeat to create enough content for truncation testing
|
||
|
||
unit_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=large_content,
|
||
context="large document test",
|
||
document_id=document_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids) > 0, "Should create units"
|
||
|
||
# Recall with very small chunk token limit to force truncation
|
||
result = await memory.recall_async(
|
||
bank_id=bank_id,
|
||
query="Tell me about the document",
|
||
budget=Budget.MID,
|
||
max_tokens=1000,
|
||
fact_type=["world"],
|
||
include_chunks=True,
|
||
max_chunk_tokens=500, # Small limit to test truncation
|
||
request_context=request_context,
|
||
)
|
||
|
||
if result.chunks:
|
||
# Check if any chunks show truncation
|
||
truncated_chunks = [
|
||
chunk_id for chunk_id, chunk_info in result.chunks.items()
|
||
if chunk_info.truncated
|
||
]
|
||
|
||
print(f"✓ Retrieved {len(result.chunks)} chunks")
|
||
if truncated_chunks:
|
||
print(f" {len(truncated_chunks)} chunks were truncated due to token limit")
|
||
else:
|
||
print(" No chunks were truncated (content within limit)")
|
||
|
||
else:
|
||
print("✓ No chunks returned (may be under token limit)")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
# ============================================================
|
||
# Memory Links Tests
|
||
# ============================================================
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_temporal_links_creation(memory, request_context):
|
||
"""
|
||
Test that temporal links are created between facts with nearby event dates.
|
||
|
||
Temporal links connect facts that occurred close in time (within 24 hours).
|
||
"""
|
||
bank_id = f"test_temporal_links_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store facts with nearby timestamps (within 24 hours)
|
||
base_date = datetime(2024, 1, 15, 10, 0, 0, tzinfo=timezone.utc)
|
||
|
||
# Fact 1 at 10:00 AM
|
||
unit_ids_1 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice started working on the authentication module.",
|
||
context="daily standup",
|
||
event_date=base_date,
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Fact 2 at 2:00 PM same day (4 hours later)
|
||
unit_ids_2 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Bob reviewed the API design document.",
|
||
context="daily standup",
|
||
event_date=base_date.replace(hour=14),
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Fact 3 at 9:00 AM next day (23 hours later)
|
||
unit_ids_3 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Charlie deployed the new database schema.",
|
||
context="daily standup",
|
||
event_date=base_date.replace(day=16, hour=9),
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids_1) > 0 and len(unit_ids_2) > 0 and len(unit_ids_3) > 0
|
||
|
||
logger.info(f"Created {len(unit_ids_1) + len(unit_ids_2) + len(unit_ids_3)} facts")
|
||
|
||
# Query the memory_links table to verify temporal links exist
|
||
async with memory._pool.acquire() as conn:
|
||
# Get all temporal links for these units
|
||
all_unit_ids = unit_ids_1 + unit_ids_2 + unit_ids_3
|
||
|
||
temporal_links = await conn.fetch(
|
||
"""
|
||
SELECT from_unit_id, to_unit_id, link_type, weight
|
||
FROM memory_links
|
||
WHERE from_unit_id::text = ANY($1)
|
||
AND link_type = 'temporal'
|
||
ORDER BY weight DESC
|
||
""",
|
||
all_unit_ids
|
||
)
|
||
|
||
logger.info(f"Found {len(temporal_links)} temporal links")
|
||
|
||
# Should have temporal links between the facts
|
||
assert len(temporal_links) > 0, "Should have created temporal links between facts with nearby dates"
|
||
|
||
# Verify link properties
|
||
for link in temporal_links:
|
||
from_id = str(link['from_unit_id'])
|
||
to_id = str(link['to_unit_id'])
|
||
logger.info(f" Link: {from_id[:8]}... -> {to_id[:8]}... (weight: {link['weight']:.2f})")
|
||
assert link['link_type'] == 'temporal', "Link type should be 'temporal'"
|
||
assert 0.0 <= link['weight'] <= 1.0, "Weight should be between 0 and 1"
|
||
|
||
logger.info("Temporal links created successfully with proper weights")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_semantic_links_creation(memory, request_context):
|
||
"""
|
||
Test that semantic links are created between facts with similar content.
|
||
|
||
Semantic links connect facts that are semantically similar based on embeddings.
|
||
"""
|
||
bank_id = f"test_semantic_links_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store facts with similar semantic content
|
||
unit_ids_1 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice is an expert in Python programming and has built many web applications.",
|
||
context="team skills",
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Similar content - should create semantic link
|
||
unit_ids_2 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Bob is proficient in Python development and specializes in building APIs.",
|
||
context="team skills",
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Different content - less likely to create strong semantic link
|
||
unit_ids_3 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="The quarterly sales meeting is scheduled for next Tuesday at 3 PM.",
|
||
context="calendar events",
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids_1) > 0 and len(unit_ids_2) > 0 and len(unit_ids_3) > 0
|
||
|
||
logger.info(f"Created {len(unit_ids_1) + len(unit_ids_2) + len(unit_ids_3)} facts")
|
||
|
||
# Query the memory_links table to verify semantic links exist
|
||
async with memory._pool.acquire() as conn:
|
||
all_unit_ids = unit_ids_1 + unit_ids_2 + unit_ids_3
|
||
|
||
semantic_links = await conn.fetch(
|
||
"""
|
||
SELECT from_unit_id, to_unit_id, link_type, weight
|
||
FROM memory_links
|
||
WHERE from_unit_id::text = ANY($1)
|
||
AND link_type = 'semantic'
|
||
ORDER BY weight DESC
|
||
""",
|
||
all_unit_ids
|
||
)
|
||
|
||
logger.info(f"Found {len(semantic_links)} semantic links")
|
||
|
||
# Should have semantic links between similar facts
|
||
assert len(semantic_links) > 0, "Should have created semantic links between similar facts"
|
||
|
||
# Verify link properties
|
||
for link in semantic_links:
|
||
from_id = str(link['from_unit_id'])
|
||
to_id = str(link['to_unit_id'])
|
||
logger.info(f" Link: {from_id[:8]}... -> {to_id[:8]}... (weight: {link['weight']:.3f})")
|
||
assert link['link_type'] == 'semantic', "Link type should be 'semantic'"
|
||
assert 0.0 <= link['weight'] <= 1.0, "Weight should be between 0 and 1"
|
||
# Semantic links typically have weight >= 0.7 (threshold)
|
||
assert link['weight'] >= 0.7, f"Semantic links should have weight >= 0.7, got {link['weight']}"
|
||
|
||
logger.info("Semantic links created successfully between similar content")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_entity_links_creation(memory, request_context):
|
||
"""
|
||
Test that entity links are created between facts that mention the same entities.
|
||
|
||
Entity links connect facts that reference the same person, place, or concept.
|
||
This is core functionality and should work consistently.
|
||
"""
|
||
bank_id = f"test_entity_links_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store facts that mention the same entities
|
||
unit_ids_1 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice joined Google as a software engineer in 2020.",
|
||
context="career history",
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Mentions same entity (Alice) - should create entity link
|
||
unit_ids_2 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice led the development of the new authentication system.",
|
||
context="project updates",
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Mentions same entity (Google) - should create entity link
|
||
unit_ids_3 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Google announced new cloud services at their annual conference.",
|
||
context="tech news",
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Different entities - no entity link expected
|
||
unit_ids_4 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Bob works at Meta on machine learning infrastructure.",
|
||
context="career history",
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids_1) > 0 and len(unit_ids_2) > 0 and len(unit_ids_3) > 0 and len(unit_ids_4) > 0
|
||
|
||
logger.info(f"Created {len(unit_ids_1) + len(unit_ids_2) + len(unit_ids_3) + len(unit_ids_4)} facts")
|
||
|
||
# Query the memory_links table to verify entity links exist
|
||
async with memory._pool.acquire() as conn:
|
||
all_unit_ids = unit_ids_1 + unit_ids_2 + unit_ids_3 + unit_ids_4
|
||
|
||
entity_links = await conn.fetch(
|
||
"""
|
||
SELECT from_unit_id, to_unit_id, link_type, weight, entity_id
|
||
FROM memory_links
|
||
WHERE from_unit_id::text = ANY($1)
|
||
AND link_type = 'entity'
|
||
ORDER BY from_unit_id, to_unit_id
|
||
""",
|
||
all_unit_ids
|
||
)
|
||
|
||
logger.info(f"Found {len(entity_links)} entity links")
|
||
|
||
# Entity extraction is core functionality and should work
|
||
assert len(entity_links) > 0, "Should have created entity links between facts with shared entities (Alice, Google)"
|
||
|
||
# Verify link properties
|
||
entities_seen = set()
|
||
for link in entity_links:
|
||
entity_id = link['entity_id']
|
||
entities_seen.add(str(entity_id))
|
||
from_id = str(link['from_unit_id'])
|
||
to_id = str(link['to_unit_id'])
|
||
logger.info(f" Link: {from_id[:8]}... -> {to_id[:8]}... via entity {str(entity_id)[:8]}...")
|
||
assert link['link_type'] == 'entity', "Link type should be 'entity'"
|
||
assert link['weight'] == 1.0, "Entity links should have weight 1.0"
|
||
assert entity_id is not None, "Entity links must reference an entity_id"
|
||
|
||
logger.info(f"Entity links created successfully for {len(entities_seen)} unique entities")
|
||
|
||
# Verify bidirectional links (entity links should be bidirectional)
|
||
link_pairs = set()
|
||
for link in entity_links:
|
||
from_id = str(link['from_unit_id'])
|
||
to_id = str(link['to_unit_id'])
|
||
entity_id = str(link['entity_id'])
|
||
link_pairs.add((from_id, to_id, entity_id))
|
||
|
||
# Check that for each (A -> B) link, there's a (B -> A) link with same entity
|
||
for from_id, to_id, entity_id in link_pairs:
|
||
reverse_exists = (to_id, from_id, entity_id) in link_pairs
|
||
assert reverse_exists, f"Entity links should be bidirectional: missing reverse link for {from_id[:8]} -> {to_id[:8]}"
|
||
|
||
logger.info("Entity links are properly bidirectional")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_people_name_extraction(memory, request_context):
|
||
"""
|
||
Test that people names are correctly extracted as entities.
|
||
|
||
This verifies that the entity resolver properly identifies and extracts
|
||
person names from content.
|
||
"""
|
||
bank_id = f"test_people_names_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store content with various people names
|
||
contents = [
|
||
"John Smith is a software engineer at Google.",
|
||
"Dr. Sarah Johnson presented her research at the conference.",
|
||
"Bob Williams and Alice Chen collaborated on the project.",
|
||
"Professor Michael Brown teaches computer science at MIT.",
|
||
]
|
||
|
||
for content in contents:
|
||
await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=content,
|
||
context="people info",
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Query entities to verify people names were extracted
|
||
async with memory._pool.acquire() as conn:
|
||
entities = await conn.fetch(
|
||
"""
|
||
SELECT canonical_name, mention_count
|
||
FROM entities
|
||
WHERE bank_id = $1
|
||
ORDER BY mention_count DESC, canonical_name
|
||
""",
|
||
bank_id
|
||
)
|
||
|
||
logger.info(f"Extracted {len(entities)} entities")
|
||
for entity in entities:
|
||
logger.info(f" - {entity['canonical_name']} (mentions: {entity['mention_count']})")
|
||
|
||
# Verify we extracted the expected people names
|
||
entity_names = {e['canonical_name'].lower() for e in entities}
|
||
|
||
# Check for expected people (names may vary slightly based on LLM extraction)
|
||
expected_people = ["john", "sarah", "bob", "alice", "michael"]
|
||
found_people = []
|
||
for person in expected_people:
|
||
matching = [name for name in entity_names if person in name]
|
||
if matching:
|
||
found_people.append(person)
|
||
logger.info(f" Found '{person}' as: {matching}")
|
||
|
||
assert len(found_people) >= 3, \
|
||
f"Should extract at least 3 people names, found: {found_people}. All entities: {entity_names}"
|
||
|
||
logger.info(f"Successfully extracted {len(found_people)} people names: {found_people}")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_mention_count_accuracy(memory, request_context):
|
||
"""
|
||
Test that mention_count is accurately tracked across retain calls.
|
||
|
||
Verifies that when an entity is mentioned multiple times across different
|
||
retain calls, the mention_count reflects the total number of mentions.
|
||
"""
|
||
bank_id = f"test_mention_count_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store content mentioning "Alice" multiple times across separate retain calls
|
||
contents = [
|
||
"Alice is a data scientist at Netflix.",
|
||
"Alice presented her research on recommendation algorithms.",
|
||
"Alice leads a team of 5 engineers.",
|
||
"Alice graduated from Stanford with honors.",
|
||
"Alice published a paper on machine learning.",
|
||
]
|
||
|
||
for content in contents:
|
||
await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=content,
|
||
context="career info",
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Check Alice's mention count
|
||
async with memory._pool.acquire() as conn:
|
||
alice_entity = await conn.fetchrow(
|
||
"""
|
||
SELECT canonical_name, mention_count
|
||
FROM entities
|
||
WHERE bank_id = $1 AND LOWER(canonical_name) LIKE '%alice%'
|
||
""",
|
||
bank_id
|
||
)
|
||
|
||
assert alice_entity is not None, "Alice entity should exist"
|
||
logger.info(f"Alice mention_count after 5 separate retains: {alice_entity['mention_count']}")
|
||
|
||
# Alice should have mention_count >= 5 (one per content item)
|
||
assert alice_entity['mention_count'] >= 5, \
|
||
f"Alice should have at least 5 mentions, got {alice_entity['mention_count']}"
|
||
|
||
logger.info(f"Mention count accuracy verified: {alice_entity['mention_count']} mentions")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_mention_count_batch_retain(memory, request_context):
|
||
"""
|
||
Test that mention_count is accurate when using batch retain with multiple items.
|
||
|
||
This specifically tests the scenario where multiple content items are retained
|
||
in a single batch call, ensuring mention_count is correctly aggregated.
|
||
"""
|
||
bank_id = f"test_mention_batch_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Batch retain with multiple items mentioning "Bob"
|
||
batch_contents = [
|
||
{"content": "Bob is a frontend developer at Microsoft.", "context": "work"},
|
||
{"content": "Bob specializes in React and TypeScript.", "context": "skills"},
|
||
{"content": "Bob has 10 years of experience.", "context": "experience"},
|
||
{"content": "Bob mentors junior developers.", "context": "mentoring"},
|
||
{"content": "Bob presented at ReactConf 2024.", "context": "conferences"},
|
||
{"content": "Bob wrote a popular open-source library.", "context": "projects"},
|
||
]
|
||
|
||
# Use retain_batch_async for batch processing
|
||
await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=batch_contents,
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Check Bob's mention count after batch retain
|
||
async with memory._pool.acquire() as conn:
|
||
bob_entity = await conn.fetchrow(
|
||
"""
|
||
SELECT canonical_name, mention_count
|
||
FROM entities
|
||
WHERE bank_id = $1 AND LOWER(canonical_name) LIKE '%bob%'
|
||
""",
|
||
bank_id
|
||
)
|
||
|
||
assert bob_entity is not None, "Bob entity should exist after batch retain"
|
||
logger.info(f"Bob mention_count after batch retain of 6 items: {bob_entity['mention_count']}")
|
||
|
||
# Bob should have mention_count >= 6 (mentioned in each batch item)
|
||
assert bob_entity['mention_count'] >= 6, \
|
||
f"Bob should have at least 6 mentions from batch retain, got {bob_entity['mention_count']}"
|
||
|
||
# Now do another batch retain with more Bob mentions
|
||
more_contents = [
|
||
{"content": "Bob loves hiking on weekends.", "context": "hobbies"},
|
||
{"content": "Bob has a dog named Max.", "context": "personal"},
|
||
]
|
||
|
||
await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=more_contents,
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Check updated mention count
|
||
async with memory._pool.acquire() as conn:
|
||
bob_entity_updated = await conn.fetchrow(
|
||
"""
|
||
SELECT canonical_name, mention_count
|
||
FROM entities
|
||
WHERE bank_id = $1 AND LOWER(canonical_name) LIKE '%bob%'
|
||
""",
|
||
bank_id
|
||
)
|
||
|
||
logger.info(f"Bob mention_count after second batch: {bob_entity_updated['mention_count']}")
|
||
|
||
# Bob should now have mention_count >= 8 (6 + 2)
|
||
assert bob_entity_updated['mention_count'] >= 8, \
|
||
f"Bob should have at least 8 mentions after second batch, got {bob_entity_updated['mention_count']}"
|
||
|
||
# Verify the increment is correct
|
||
increment = bob_entity_updated['mention_count'] - bob_entity['mention_count']
|
||
assert increment >= 2, \
|
||
f"Mention count should have increased by at least 2, but increased by {increment}"
|
||
|
||
logger.info(f"Batch retain mention count verified: {bob_entity['mention_count']} -> {bob_entity_updated['mention_count']}")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_causal_links_creation(memory, request_context):
|
||
"""
|
||
Test that causal links are created between facts with causal relationships.
|
||
|
||
Causal links connect facts where one causes, enables, or prevents another.
|
||
Note: This depends on LLM extracting causal relationships, which may be non-deterministic.
|
||
"""
|
||
bank_id = f"test_causal_links_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store content with explicit causal relationships
|
||
# Using clear cause-and-effect language to maximize LLM detection
|
||
content = """
|
||
Alice completed the authentication module on Monday. Because Alice finished the auth module,
|
||
Bob was able to start integrating it with the API on Tuesday. Bob's API integration enabled
|
||
Charlie to begin testing the complete user flow on Wednesday. The successful testing caused
|
||
the team to schedule the production deployment for Friday.
|
||
"""
|
||
|
||
unit_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=content,
|
||
context="project timeline",
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids) > 0, "Should have created facts"
|
||
logger.info(f"Created {len(unit_ids)} facts from causal content")
|
||
|
||
# Query the memory_links table to check for causal links
|
||
async with memory._pool.acquire() as conn:
|
||
causal_links = await conn.fetch(
|
||
"""
|
||
SELECT from_unit_id, to_unit_id, link_type, weight
|
||
FROM memory_links
|
||
WHERE from_unit_id::text = ANY($1)
|
||
AND link_type IN ('causes', 'caused_by', 'enables', 'prevents')
|
||
ORDER BY link_type, weight DESC
|
||
""",
|
||
unit_ids
|
||
)
|
||
|
||
logger.info(f"Found {len(causal_links)} causal links")
|
||
|
||
if len(causal_links) > 0:
|
||
# Verify link properties
|
||
causal_types = {}
|
||
for link in causal_links:
|
||
link_type = link['link_type']
|
||
causal_types[link_type] = causal_types.get(link_type, 0) + 1
|
||
from_id = str(link['from_unit_id'])
|
||
to_id = str(link['to_unit_id'])
|
||
logger.info(f" Link: {from_id[:8]}... -> {to_id[:8]}... ({link_type}, weight: {link['weight']:.2f})")
|
||
assert link['link_type'] in ['causes', 'caused_by', 'enables', 'prevents'], \
|
||
f"Causal link type must be valid, got '{link['link_type']}'"
|
||
assert 0.0 <= link['weight'] <= 1.0, "Weight should be between 0 and 1"
|
||
|
||
logger.info("Causal links created successfully:")
|
||
for link_type, count in causal_types.items():
|
||
logger.info(f" - {link_type}: {count} links")
|
||
else:
|
||
logger.warning("No causal links detected (LLM may not have extracted causal relationships)")
|
||
logger.info(" This is expected as causal extraction depends on LLM interpretation")
|
||
|
||
# This test passes even if no causal links are found, since causal extraction
|
||
# is non-deterministic and depends on LLM behavior
|
||
logger.info("Test completed (causal link extraction is LLM-dependent)")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_all_link_types_together(memory, request_context):
|
||
"""
|
||
Integration test: Verify all link types can be created in a single retain operation.
|
||
|
||
Tests that temporal, semantic, entity, and potentially causal links are all
|
||
created when appropriate conditions are met.
|
||
"""
|
||
bank_id = f"test_all_links_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store multiple related facts that should trigger all link types
|
||
base_date = datetime(2024, 1, 15, 10, 0, 0, tzinfo=timezone.utc)
|
||
|
||
# Fact 1: Alice at time T
|
||
unit_ids_1 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice completed the Python backend service for the authentication system.",
|
||
context="sprint review",
|
||
event_date=base_date,
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Fact 2: Related to Alice, similar topic (Python), close in time
|
||
unit_ids_2 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice optimized the Python code and improved the authentication performance by 40%.",
|
||
context="sprint review",
|
||
event_date=base_date.replace(hour=14), # Same day, 4 hours later
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Fact 3: Related to Alice, different topic but same entity
|
||
unit_ids_3 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice presented the security architecture at the team meeting.",
|
||
context="team meeting",
|
||
event_date=base_date.replace(day=16), # Next day
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids_1) > 0 and len(unit_ids_2) > 0 and len(unit_ids_3) > 0
|
||
|
||
logger.info(f"Created {len(unit_ids_1) + len(unit_ids_2) + len(unit_ids_3)} facts")
|
||
|
||
# Query for all link types
|
||
async with memory._pool.acquire() as conn:
|
||
all_unit_ids = unit_ids_1 + unit_ids_2 + unit_ids_3
|
||
|
||
all_links = await conn.fetch(
|
||
"""
|
||
SELECT link_type, COUNT(*) as count
|
||
FROM memory_links
|
||
WHERE from_unit_id::text = ANY($1)
|
||
GROUP BY link_type
|
||
ORDER BY link_type
|
||
""",
|
||
all_unit_ids
|
||
)
|
||
|
||
logger.info("Link types created:")
|
||
link_types_found = {}
|
||
for row in all_links:
|
||
link_type = row['link_type']
|
||
count = row['count']
|
||
link_types_found[link_type] = count
|
||
logger.info(f" - {link_type}: {count} links")
|
||
|
||
# Should have temporal, semantic, and entity links
|
||
assert 'temporal' in link_types_found, "Should have temporal links (facts with nearby dates)"
|
||
assert 'semantic' in link_types_found, "Should have semantic links (similar content about Python/auth)"
|
||
assert 'entity' in link_types_found, "Should have entity links (all mention Alice)"
|
||
|
||
logger.info(f"Successfully created {len(link_types_found)} different link types")
|
||
logger.info("All major link types (temporal, semantic, entity) are working correctly")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_semantic_links_within_same_batch(memory, request_context):
|
||
"""
|
||
Test that semantic links are created between facts retained in the SAME batch.
|
||
|
||
This is a regression test - semantic links should connect similar facts
|
||
even when they are retained together in a single call.
|
||
"""
|
||
bank_id = f"test_semantic_batch_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Retain multiple semantically similar facts in ONE batch
|
||
contents = [
|
||
{"content": "Alice is an expert in Python programming and machine learning.", "context": "team skills"},
|
||
{"content": "Bob specializes in Python development and data science.", "context": "team skills"},
|
||
{"content": "Charlie works with Python for backend API development.", "context": "team skills"},
|
||
]
|
||
|
||
result = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=contents,
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Flatten the list of lists
|
||
unit_ids = [uid for sublist in result for uid in sublist]
|
||
|
||
assert len(unit_ids) >= 3, f"Should have created at least 3 facts, got {len(unit_ids)}"
|
||
logger.info(f"Created {len(unit_ids)} facts in single batch")
|
||
|
||
# Query semantic links between these units
|
||
async with memory._pool.acquire() as conn:
|
||
semantic_links = await conn.fetch(
|
||
"""
|
||
SELECT from_unit_id, to_unit_id, weight
|
||
FROM memory_links
|
||
WHERE from_unit_id::text = ANY($1)
|
||
AND to_unit_id::text = ANY($1)
|
||
AND link_type = 'semantic'
|
||
""",
|
||
unit_ids
|
||
)
|
||
|
||
logger.info(f"Found {len(semantic_links)} semantic links within the batch")
|
||
|
||
# All three facts mention Python - they should be linked to each other
|
||
assert len(semantic_links) > 0, (
|
||
"REGRESSION: Semantic links should be created between similar facts "
|
||
"retained in the same batch, but none were found"
|
||
)
|
||
|
||
# Log the links for debugging
|
||
for link in semantic_links:
|
||
logger.info(f" Semantic link: {str(link['from_unit_id'])[:8]}... -> {str(link['to_unit_id'])[:8]}... (weight: {link['weight']:.3f})")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_semantic_links_phase1_ann_cross_batch(memory, request_context):
|
||
"""
|
||
Test that Phase 1 ANN search creates semantic links between facts from
|
||
DIFFERENT retain batches.
|
||
|
||
The semantic ANN search runs in Phase 1 on a separate connection (outside
|
||
the write transaction) using placeholder unit IDs to avoid TimeoutErrors
|
||
from HNSW index contention under concurrent load. This test verifies that:
|
||
1. Phase 1 ANN with placeholder IDs works correctly
|
||
2. Placeholder IDs are remapped to real unit IDs before insertion
|
||
3. Cross-batch semantic links are created between similar facts
|
||
"""
|
||
bank_id = f"test_semantic_phase1_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# First batch: store some facts about Python
|
||
await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Alice is an expert Python developer who builds web applications using FastAPI.",
|
||
context="team skills",
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Second batch: store similar facts — Phase 1 ANN should find the first batch's
|
||
# facts via HNSW index and create cross-batch semantic links
|
||
unit_ids_2 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Bob specializes in Python programming and creates REST APIs with FastAPI.",
|
||
context="team skills",
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(unit_ids_2) > 0
|
||
|
||
# Verify cross-batch semantic links exist
|
||
async with memory._pool.acquire() as conn:
|
||
cross_batch_links = await conn.fetch(
|
||
"""
|
||
SELECT from_unit_id, to_unit_id, weight
|
||
FROM memory_links
|
||
WHERE from_unit_id::text = ANY($1)
|
||
AND link_type = 'semantic'
|
||
AND to_unit_id::text != ALL($1)
|
||
""",
|
||
unit_ids_2,
|
||
)
|
||
|
||
logger.info(f"Cross-batch semantic links from batch 2: {len(cross_batch_links)}")
|
||
for link in cross_batch_links:
|
||
logger.info(
|
||
f" {str(link['from_unit_id'])[:8]}... -> {str(link['to_unit_id'])[:8]}... "
|
||
f"(weight: {link['weight']:.3f})"
|
||
)
|
||
|
||
# Phase 1 ANN should have found similar facts from batch 1
|
||
assert len(cross_batch_links) > 0, (
|
||
"Phase 1 ANN search should create semantic links between similar facts "
|
||
"from different retain batches. This tests that placeholder unit IDs are "
|
||
"correctly remapped to real IDs after insert_facts_batch."
|
||
)
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_temporal_links_within_same_batch(memory, request_context):
|
||
"""
|
||
Test that temporal links are created between facts retained in the SAME batch.
|
||
|
||
This is a regression test - temporal links should connect facts with nearby
|
||
event dates even when they are retained together in a single call.
|
||
"""
|
||
bank_id = f"test_temporal_batch_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Retain multiple facts with nearby timestamps in ONE batch
|
||
base_date = datetime(2024, 6, 15, 10, 0, 0, tzinfo=timezone.utc)
|
||
|
||
contents = [
|
||
{
|
||
"content": "Morning standup: Alice presented the sprint goals.",
|
||
"context": "daily meeting",
|
||
"event_date": base_date
|
||
},
|
||
{
|
||
"content": "Bob demoed the new feature after standup.",
|
||
"context": "daily meeting",
|
||
"event_date": base_date + timedelta(hours=1) # 1 hour later
|
||
},
|
||
{
|
||
"content": "Charlie reviewed the pull requests in the afternoon.",
|
||
"context": "daily meeting",
|
||
"event_date": base_date + timedelta(hours=4) # 4 hours later
|
||
},
|
||
]
|
||
|
||
result = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=contents,
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Flatten the list of lists
|
||
unit_ids = [uid for sublist in result for uid in sublist]
|
||
|
||
assert len(unit_ids) >= 3, f"Should have created at least 3 facts, got {len(unit_ids)}"
|
||
logger.info(f"Created {len(unit_ids)} facts in single batch")
|
||
|
||
# Query temporal links between these units
|
||
async with memory._pool.acquire() as conn:
|
||
temporal_links = await conn.fetch(
|
||
"""
|
||
SELECT from_unit_id, to_unit_id, weight
|
||
FROM memory_links
|
||
WHERE from_unit_id::text = ANY($1)
|
||
AND to_unit_id::text = ANY($1)
|
||
AND link_type = 'temporal'
|
||
""",
|
||
unit_ids
|
||
)
|
||
|
||
logger.info(f"Found {len(temporal_links)} temporal links within the batch")
|
||
|
||
# All three facts are within 24 hours - they should be linked to each other
|
||
assert len(temporal_links) > 0, (
|
||
"REGRESSION: Temporal links should be created between facts with nearby dates "
|
||
"retained in the same batch, but none were found"
|
||
)
|
||
|
||
# Log the links for debugging
|
||
for link in temporal_links:
|
||
logger.info(f" Temporal link: {str(link['from_unit_id'])[:8]}... -> {str(link['to_unit_id'])[:8]}... (weight: {link['weight']:.3f})")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_user_provided_entities(memory, request_context):
|
||
"""
|
||
Test that user-provided entities are merged with auto-extracted entities.
|
||
|
||
This tests the feature added in PR #91 where users can provide entities
|
||
via the 'entities' field in the retain request. These should be combined
|
||
with LLM-extracted entities, with case-insensitive deduplication.
|
||
"""
|
||
bank_id = f"test_user_entities_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Store content with user-provided entities
|
||
# The content mentions "Alice" which LLM might extract,
|
||
# but we also provide "ProjectX" and "ACME Corp" which may not be in the text
|
||
contents = [
|
||
{
|
||
"content": "Alice completed the quarterly report.",
|
||
"context": "work update",
|
||
"entities": [
|
||
{"text": "ProjectX", "type": "PROJECT"},
|
||
{"text": "ACME Corp", "type": "ORG"},
|
||
{"text": "Alice"}, # May also be extracted by LLM (dedup test)
|
||
],
|
||
}
|
||
]
|
||
|
||
result = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=contents,
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Flatten the list of lists
|
||
unit_ids = [uid for sublist in result for uid in sublist]
|
||
assert len(unit_ids) > 0, "Should have created at least one fact"
|
||
|
||
logger.info(f"Created {len(unit_ids)} facts with user-provided entities")
|
||
|
||
# Query entity links to verify user-provided entities were stored
|
||
async with memory._pool.acquire() as conn:
|
||
# Get all entities linked to our facts via the unit_entities junction table
|
||
entity_rows = await conn.fetch(
|
||
"""
|
||
SELECT DISTINCT e.canonical_name
|
||
FROM entities e
|
||
JOIN unit_entities ue ON e.id = ue.entity_id
|
||
WHERE ue.unit_id::text = ANY($1)
|
||
""",
|
||
unit_ids
|
||
)
|
||
|
||
entity_names = {row['canonical_name'].lower() for row in entity_rows}
|
||
logger.info(f"Found entities linked to facts: {[row['canonical_name'] for row in entity_rows]}")
|
||
|
||
# Verify user-provided entities are present
|
||
assert "projectx" in entity_names, "User-provided entity 'ProjectX' should be linked"
|
||
assert "acme corp" in entity_names, "User-provided entity 'ACME Corp' should be linked"
|
||
|
||
# Alice should be present (either from LLM extraction or user-provided)
|
||
assert "alice" in entity_names, "Entity 'Alice' should be linked"
|
||
|
||
logger.info("✓ User-provided entities successfully merged with extracted entities")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
def test_recall_result_model_empty_construction():
|
||
"""
|
||
Test that RecallResultModel can be constructed with empty results.
|
||
|
||
This is a regression test for the bug where constructing an empty RecallResultModel
|
||
would cause an UnboundLocalError because RecallResult was imported as RecallResultModel
|
||
but the code mistakenly used the wrong name.
|
||
|
||
The fix ensures RecallResultModel is used consistently throughout memory_engine.py.
|
||
"""
|
||
from hindsight_api.engine.response_models import RecallResult
|
||
|
||
# This should not raise any errors
|
||
result = RecallResult(results=[], entities={}, chunks={})
|
||
|
||
assert result is not None, "Should create a result object"
|
||
assert result.results == [], "Should have empty results"
|
||
assert result.entities == {}, "Should have empty entities"
|
||
assert result.chunks == {}, "Should have empty chunks"
|
||
|
||
logger.info("✓ RecallResult empty construction works correctly")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_custom_extraction_mode():
|
||
"""
|
||
Test that custom extraction mode uses custom guidelines from env variable.
|
||
|
||
This test verifies that when HINDSIGHT_API_RETAIN_EXTRACTION_MODE=custom and
|
||
HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS is set, the fact extraction uses the
|
||
custom guidelines while keeping structural parts intact.
|
||
"""
|
||
import os
|
||
|
||
from hindsight_api import LLMConfig
|
||
from hindsight_api.config import _get_raw_config, clear_config_cache
|
||
from hindsight_api.engine.retain.fact_extraction import extract_facts_from_text
|
||
|
||
# Save original env vars
|
||
original_mode = os.getenv("HINDSIGHT_API_RETAIN_EXTRACTION_MODE")
|
||
original_instructions = os.getenv("HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS")
|
||
|
||
try:
|
||
# Set custom extraction mode with challenging language-specific guidelines
|
||
os.environ["HINDSIGHT_API_RETAIN_EXTRACTION_MODE"] = "custom"
|
||
os.environ["HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS"] = """ONLY extract facts that are in ITALIAN language.
|
||
|
||
DO NOT extract:
|
||
❌ Facts in English
|
||
❌ Facts in any other language besides Italian
|
||
|
||
If the text contains both Italian and English content, extract ONLY the Italian facts."""
|
||
|
||
# Clear config cache to pick up new env vars
|
||
clear_config_cache()
|
||
|
||
# Test content with BOTH Italian (should extract) and English (should NOT extract) facts
|
||
# This is a much harder test than filtering greetings
|
||
text = """
|
||
The team discussed the new architecture. We will use microservices.
|
||
|
||
Il database PostgreSQL ha ridotto la latenza delle query del 60%.
|
||
Alice ha suggerito di usare il connection pooling per migliorare le prestazioni.
|
||
|
||
Bob mentioned that the API endpoint is ready for testing.
|
||
The deployment pipeline has been updated to use Kubernetes.
|
||
|
||
Marco ha completato la revisione del codice e ha approvato le modifiche.
|
||
Il sistema di autenticazione è stato migrato a OAuth 2.0.
|
||
"""
|
||
|
||
llm_config = LLMConfig.from_env()
|
||
|
||
facts, _, _ = await extract_facts_from_text(
|
||
text=text,
|
||
event_date=datetime(2024, 1, 15, tzinfo=timezone.utc),
|
||
context="team meeting notes",
|
||
llm_config=llm_config,
|
||
agent_name="TestUser",
|
||
config=_get_raw_config(),
|
||
)
|
||
|
||
logger.info(f"\nExtracted {len(facts)} facts with custom mode (Italian only):")
|
||
for i, fact in enumerate(facts):
|
||
logger.info(f" {i+1}. {fact.fact}")
|
||
|
||
assert len(facts) > 0, "Should extract at least one Italian fact"
|
||
|
||
# All facts text
|
||
all_facts_text = " ".join([f.fact for f in facts])
|
||
|
||
# Should HAVE Italian content
|
||
italian_keywords = ["postgresql", "latenza", "query", "alice", "connection pooling", "prestazioni",
|
||
"marco", "revisione", "codice", "autenticazione", "oauth"]
|
||
has_italian = any(keyword in all_facts_text.lower() for keyword in italian_keywords)
|
||
assert has_italian, f"Should extract Italian facts. Got: {all_facts_text}"
|
||
|
||
# Should NOT have English-only content
|
||
# These are facts that appear ONLY in English sections
|
||
english_only_keywords = ["microservices", "bob", "api endpoint", "testing", "deployment pipeline", "kubernetes"]
|
||
|
||
# Check if facts contain English-only content (this would be wrong)
|
||
facts_lower = all_facts_text.lower()
|
||
found_english_only = [kw for kw in english_only_keywords if kw in facts_lower]
|
||
|
||
if found_english_only:
|
||
logger.warning(f"⚠ Found English-only keywords in facts: {found_english_only}")
|
||
logger.warning(f" Facts: {all_facts_text}")
|
||
logger.warning(" This may indicate the LLM is not strictly following language-specific custom guidelines")
|
||
# Log but don't fail - LLM behavior can vary
|
||
else:
|
||
logger.info("✓ Successfully extracted only Italian facts, ignored English facts")
|
||
|
||
# At least verify we have some Italian indicators
|
||
italian_indicators = ["latenza", "prestazioni", "revisione", "codice", "autenticazione"]
|
||
italian_count = sum(1 for ind in italian_indicators if ind in facts_lower)
|
||
|
||
assert italian_count >= 1, \
|
||
f"Should extract facts with Italian words. Found {italian_count} Italian indicators in: {all_facts_text}"
|
||
|
||
logger.info("✓ Custom extraction mode works with language-specific guidelines")
|
||
logger.info(f"✓ Extracted {len(facts)} Italian facts, found {italian_count} Italian indicators")
|
||
|
||
finally:
|
||
# Restore original env vars
|
||
if original_mode is not None:
|
||
os.environ["HINDSIGHT_API_RETAIN_EXTRACTION_MODE"] = original_mode
|
||
else:
|
||
os.environ.pop("HINDSIGHT_API_RETAIN_EXTRACTION_MODE", None)
|
||
|
||
if original_instructions is not None:
|
||
os.environ["HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS"] = original_instructions
|
||
else:
|
||
os.environ.pop("HINDSIGHT_API_RETAIN_CUSTOM_INSTRUCTIONS", None)
|
||
|
||
# Clear cache again to restore original config
|
||
clear_config_cache()
|
||
|
||
|
||
def test_apply_strategy():
|
||
"""
|
||
Unit test for apply_strategy:
|
||
- Known strategy applies overrides on top of resolved config
|
||
- Unknown strategy returns config unchanged with a warning
|
||
- Non-hierarchical fields in a strategy are silently ignored
|
||
- entity_labels and entities_allow_free_form are overridable
|
||
"""
|
||
from hindsight_api.config import _get_raw_config, clear_config_cache
|
||
from hindsight_api.config_resolver import apply_strategy
|
||
|
||
clear_config_cache()
|
||
base_config = _get_raw_config()
|
||
|
||
strategies = {
|
||
"documents": {
|
||
"retain_extraction_mode": "chunks",
|
||
"retain_chunk_size": 800,
|
||
"entities_allow_free_form": False,
|
||
},
|
||
"bad_field": {
|
||
"database_url": "should-be-ignored", # static field, not hierarchical
|
||
"retain_extraction_mode": "verbose",
|
||
},
|
||
}
|
||
config_with_strategies = base_config.__class__(
|
||
**{**base_config.__dict__, "retain_strategies": strategies}
|
||
)
|
||
|
||
# Known strategy: overrides applied
|
||
result = apply_strategy(config_with_strategies, "documents")
|
||
assert result.retain_extraction_mode == "chunks"
|
||
assert result.retain_chunk_size == 800
|
||
assert result.entities_allow_free_form is False
|
||
|
||
# Non-hierarchical field silently ignored, hierarchical one applied
|
||
result2 = apply_strategy(config_with_strategies, "bad_field")
|
||
assert result2.retain_extraction_mode == "verbose"
|
||
assert result2.database_url == base_config.database_url # unchanged
|
||
|
||
# Unknown strategy: config returned unchanged
|
||
result3 = apply_strategy(config_with_strategies, "nonexistent")
|
||
assert result3.retain_extraction_mode == base_config.retain_extraction_mode
|
||
|
||
|
||
def test_collapse_to_verbatim_single_fact_per_chunk():
|
||
"""
|
||
Unit test for _collapse_to_verbatim:
|
||
- One fact per chunk → text overridden with original chunk text
|
||
- Two facts from same chunk → collapsed to one, entities merged
|
||
"""
|
||
from hindsight_api.engine.retain.fact_extraction import _collapse_to_verbatim
|
||
from hindsight_api.engine.retain.types import ChunkMetadata, ExtractedFact
|
||
|
||
chunks = [
|
||
ChunkMetadata(chunk_text="Alice went to Paris.", fact_count=1, content_index=0, chunk_index=0),
|
||
ChunkMetadata(chunk_text="Bob fixed the bug yesterday.", fact_count=2, content_index=0, chunk_index=1),
|
||
]
|
||
|
||
facts = [
|
||
ExtractedFact(fact_text="LLM paraphrase of Alice in Paris", fact_type="world", entities=["Alice", "Paris"], chunk_index=0, content_index=0),
|
||
ExtractedFact(fact_text="LLM first fact about Bob", fact_type="world", entities=["Bob"], chunk_index=1, content_index=0),
|
||
ExtractedFact(fact_text="LLM second fact about bug", fact_type="world", entities=["bug"], chunk_index=1, content_index=0),
|
||
]
|
||
|
||
result = _collapse_to_verbatim(facts, chunks)
|
||
|
||
assert len(result) == 2, "Should produce exactly one fact per chunk"
|
||
|
||
# Chunk 0: text overridden with original chunk text
|
||
assert result[0].fact_text == "Alice went to Paris.", "Text must be the raw chunk text"
|
||
assert result[0].entities == ["Alice", "Paris"]
|
||
|
||
# Chunk 1: collapsed to one fact, entities merged from both LLM facts
|
||
assert result[1].fact_text == "Bob fixed the bug yesterday.", "Text must be the raw chunk text"
|
||
assert "Bob" in result[1].entities
|
||
assert "bug" in result[1].entities
|
||
|
||
|
||
def test_chunks_extraction_mode():
|
||
"""
|
||
Unit test for chunks mode: no LLM, chunks stored as-is, zero token usage.
|
||
"""
|
||
import asyncio
|
||
import os
|
||
|
||
from hindsight_api.config import _get_raw_config, clear_config_cache
|
||
from hindsight_api.engine.retain.fact_extraction import extract_facts_from_contents
|
||
from hindsight_api.engine.retain.types import RetainContent
|
||
|
||
original_mode = os.getenv("HINDSIGHT_API_RETAIN_EXTRACTION_MODE")
|
||
|
||
try:
|
||
os.environ["HINDSIGHT_API_RETAIN_EXTRACTION_MODE"] = "chunks"
|
||
clear_config_cache()
|
||
|
||
contents = [
|
||
RetainContent(
|
||
content="Alice joined the infrastructure team on March 5, 2024.",
|
||
event_date=datetime(2024, 3, 10, tzinfo=timezone.utc),
|
||
entities=[{"text": "Alice"}, {"text": "infrastructure team"}],
|
||
),
|
||
RetainContent(content="Bob fixed the critical bug in the payment service."),
|
||
]
|
||
|
||
facts, chunks, usage = asyncio.get_event_loop().run_until_complete(
|
||
extract_facts_from_contents(
|
||
contents=contents,
|
||
llm_config=None, # Must not be called
|
||
agent_name="TestAgent",
|
||
config=_get_raw_config(),
|
||
)
|
||
)
|
||
|
||
# One fact per chunk (both contents fit in one chunk each)
|
||
assert len(facts) == len(chunks) == 2
|
||
|
||
# Text preserved exactly
|
||
assert facts[0].fact_text == contents[0].content
|
||
assert facts[1].fact_text == contents[1].content
|
||
|
||
# No LLM-extracted entities (user-provided entities handled downstream)
|
||
assert facts[0].entities == []
|
||
assert facts[1].entities == []
|
||
|
||
# Zero token usage
|
||
assert usage.total_tokens == 0
|
||
|
||
logger.info("✓ chunks mode: no LLM call, chunks stored as-is, zero token usage")
|
||
|
||
finally:
|
||
if original_mode is not None:
|
||
os.environ["HINDSIGHT_API_RETAIN_EXTRACTION_MODE"] = original_mode
|
||
else:
|
||
os.environ.pop("HINDSIGHT_API_RETAIN_EXTRACTION_MODE", None)
|
||
clear_config_cache()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_verbatim_extraction_mode():
|
||
"""
|
||
Integration test for verbatim extraction mode.
|
||
|
||
Verifies that:
|
||
1. Each chunk produces exactly one fact
|
||
2. The fact text is the original chunk text, not a paraphrase
|
||
3. Entities are still extracted by the LLM
|
||
4. Temporal info (occurred_start) is still extracted
|
||
"""
|
||
import os
|
||
|
||
from hindsight_api import LLMConfig
|
||
from hindsight_api.config import _get_raw_config, clear_config_cache
|
||
from hindsight_api.engine.retain.fact_extraction import extract_facts_from_contents
|
||
from hindsight_api.engine.retain.types import RetainContent
|
||
|
||
original_mode = os.getenv("HINDSIGHT_API_RETAIN_EXTRACTION_MODE")
|
||
|
||
try:
|
||
os.environ["HINDSIGHT_API_RETAIN_EXTRACTION_MODE"] = "verbatim"
|
||
clear_config_cache()
|
||
|
||
text = (
|
||
"Alice joined the infrastructure team on March 5, 2024. "
|
||
"She holds a CKA certification and has 5 years of Kubernetes experience."
|
||
)
|
||
|
||
llm_config = LLMConfig.from_env()
|
||
contents = [RetainContent(content=text, event_date=datetime(2024, 3, 10, tzinfo=timezone.utc), context="onboarding notes")]
|
||
facts, chunks, _ = await extract_facts_from_contents(
|
||
contents=contents,
|
||
llm_config=llm_config,
|
||
agent_name="TestAgent",
|
||
config=_get_raw_config(),
|
||
)
|
||
|
||
logger.info(f"Verbatim mode extracted {len(facts)} facts from {len(chunks)} chunks")
|
||
for i, f in enumerate(facts):
|
||
logger.info(f" fact[{i}]: {f.fact_text!r} entities={f.entities}")
|
||
|
||
# One fact per chunk
|
||
assert len(facts) == len(chunks), "Verbatim mode must produce exactly one fact per chunk"
|
||
|
||
# Text must match the original chunk exactly
|
||
for fact, chunk in zip(facts, chunks):
|
||
assert fact.fact_text == chunk.chunk_text, (
|
||
f"fact_text must equal original chunk text.\n"
|
||
f" expected: {chunk.chunk_text!r}\n"
|
||
f" got: {fact.fact_text!r}"
|
||
)
|
||
|
||
# Entities should still be extracted
|
||
all_entities = [e for f in facts for e in f.entities]
|
||
assert any("alice" in e.lower() for e in all_entities), (
|
||
f"Expected entity 'Alice' to be extracted. Entities: {all_entities}"
|
||
)
|
||
|
||
logger.info("✓ Verbatim mode preserves chunk text and still extracts entities")
|
||
|
||
finally:
|
||
if original_mode is not None:
|
||
os.environ["HINDSIGHT_API_RETAIN_EXTRACTION_MODE"] = original_mode
|
||
else:
|
||
os.environ.pop("HINDSIGHT_API_RETAIN_EXTRACTION_MODE", None)
|
||
clear_config_cache()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_retain_batch_with_per_item_tags_on_document(memory, request_context):
|
||
"""
|
||
Test that per-item tags are correctly stored on documents.
|
||
|
||
This test verifies the fix for a bug where per-item tags in content dictionaries
|
||
were not being merged and passed to document tracking, causing tags to be lost
|
||
even though they were correctly sent through the API.
|
||
|
||
Without the fix, this test would fail because:
|
||
- Tags are correctly passed in the content dict
|
||
- Tags are correctly stored on memory_units (facts)
|
||
- BUT tags were NOT stored on the document record itself
|
||
"""
|
||
bank_id = f"test_doc_tags_{datetime.now(timezone.utc).timestamp()}"
|
||
document_id = "app-state-testuser"
|
||
|
||
try:
|
||
# Retain content with per-item tags (simulating the TasteAI use case)
|
||
contents = [
|
||
{
|
||
"content": '{"username":"testuser","meals":[],"preferences":{"nickname":"testuser"}}',
|
||
"document_id": document_id,
|
||
"tags": ["user:testuser", "app-type:taste-ai"],
|
||
}
|
||
]
|
||
|
||
result = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=contents,
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert len(result) > 0, "Should have retained content"
|
||
print("\n=== Retained content with tags ===")
|
||
|
||
# Retrieve the document
|
||
doc = await memory.get_document(
|
||
document_id=document_id,
|
||
bank_id=bank_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
assert doc is not None, "Document should exist"
|
||
assert "tags" in doc, "Document should have tags field"
|
||
|
||
# This is the critical assertion - tags should be stored on the document
|
||
doc_tags = doc["tags"] or []
|
||
print(f"Document tags: {doc_tags}")
|
||
|
||
assert "user:testuser" in doc_tags, \
|
||
f"Document should have 'user:testuser' tag, but got: {doc_tags}"
|
||
assert "app-type:taste-ai" in doc_tags, \
|
||
f"Document should have 'app-type:taste-ai' tag, but got: {doc_tags}"
|
||
|
||
print("✓ Per-item tags correctly stored on document")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
print(f"\n=== Cleaned up bank: {bank_id} ===")
|
||
|
||
|
||
def test_retain_mission_injected_into_prompt():
|
||
"""Test that retain_mission is injected as a FOCUS section into any extraction mode."""
|
||
from unittest.mock import MagicMock
|
||
|
||
from hindsight_api.engine.retain.fact_extraction import _build_extraction_prompt_and_schema
|
||
|
||
spec = "Focus on technical decisions and architecture choices only."
|
||
|
||
# Test with concise mode
|
||
config = MagicMock()
|
||
config.retain_extraction_mode = "concise"
|
||
config.retain_mission = spec
|
||
config.retain_custom_instructions = None
|
||
config.retain_extract_causal_links = False
|
||
|
||
prompt, _ = _build_extraction_prompt_and_schema(config)
|
||
assert spec in prompt
|
||
assert "FOCUS" in prompt
|
||
|
||
# retain_mission is injected into verbose mode as well
|
||
config.retain_extraction_mode = "verbose"
|
||
prompt_verbose, _ = _build_extraction_prompt_and_schema(config)
|
||
assert spec in prompt_verbose
|
||
assert "FOCUS" in prompt_verbose
|
||
|
||
|
||
def test_retain_mission_absent_when_not_set():
|
||
"""Test that no FOCUS section appears when retain_mission is not set."""
|
||
from unittest.mock import MagicMock
|
||
|
||
from hindsight_api.engine.retain.fact_extraction import _build_extraction_prompt_and_schema
|
||
|
||
config = MagicMock()
|
||
config.retain_extraction_mode = "concise"
|
||
config.retain_mission = None
|
||
config.retain_custom_instructions = None
|
||
config.retain_extract_causal_links = False
|
||
|
||
prompt, _ = _build_extraction_prompt_and_schema(config)
|
||
assert "FOCUS" not in prompt
|
||
assert "retain_mission_section" not in prompt
|
||
|
||
|
||
def test_retain_mission_config_loaded_from_env():
|
||
"""Test that retain_mission is loaded from env and is a configurable field."""
|
||
import os
|
||
|
||
from hindsight_api.config import HindsightConfig, _get_raw_config, clear_config_cache
|
||
|
||
original = os.getenv("HINDSIGHT_API_RETAIN_MISSION")
|
||
try:
|
||
os.environ["HINDSIGHT_API_RETAIN_MISSION"] = "Only technical decisions."
|
||
clear_config_cache()
|
||
config = _get_raw_config()
|
||
assert config.retain_mission == "Only technical decisions."
|
||
assert "retain_mission" in HindsightConfig.get_configurable_fields()
|
||
finally:
|
||
if original is None:
|
||
os.environ.pop("HINDSIGHT_API_RETAIN_MISSION", None)
|
||
else:
|
||
os.environ["HINDSIGHT_API_RETAIN_MISSION"] = original
|
||
clear_config_cache()
|
||
|
||
|
||
def test_strategy_overrides_extraction_mode_for_chunks():
|
||
"""
|
||
Unit test: a named strategy with retain_extraction_mode=chunks causes
|
||
extract_facts_from_contents to skip the LLM and return verbatim chunks.
|
||
"""
|
||
import asyncio
|
||
|
||
from hindsight_api.config import _get_raw_config, clear_config_cache
|
||
from hindsight_api.config_resolver import apply_strategy
|
||
from hindsight_api.engine.retain.fact_extraction import extract_facts_from_contents
|
||
from hindsight_api.engine.retain.types import RetainContent
|
||
|
||
clear_config_cache()
|
||
base_config = _get_raw_config()
|
||
|
||
# Build a config that has a strategy overriding to chunks
|
||
strategies = {"fast": {"retain_extraction_mode": "chunks"}}
|
||
config_with_strategies = base_config.__class__(
|
||
**{**base_config.__dict__, "retain_strategies": strategies}
|
||
)
|
||
strategy_config = apply_strategy(config_with_strategies, "fast")
|
||
assert strategy_config.retain_extraction_mode == "chunks"
|
||
|
||
contents = [
|
||
RetainContent(content="Alice deployed the new API on Monday."),
|
||
RetainContent(content="Bob reviewed the pull request."),
|
||
]
|
||
|
||
facts, chunks, usage = asyncio.get_event_loop().run_until_complete(
|
||
extract_facts_from_contents(
|
||
contents=contents,
|
||
llm_config=None, # chunks must not call the LLM
|
||
agent_name="TestAgent",
|
||
config=strategy_config,
|
||
)
|
||
)
|
||
|
||
assert len(facts) == 2
|
||
assert facts[0].fact_text == contents[0].content
|
||
assert facts[1].fact_text == contents[1].content
|
||
assert usage.total_tokens == 0
|
||
logger.info("✓ strategy with chunks mode: no LLM, verbatim chunks, zero tokens")
|
||
|
||
|
||
def test_retain_request_per_item_strategy_field():
|
||
"""
|
||
Unit test: MemoryItem accepts a strategy field; items with different strategies
|
||
are grouped correctly by per-item strategy.
|
||
"""
|
||
from hindsight_api.api.http import RetainRequest
|
||
|
||
request = RetainRequest.model_validate(
|
||
{
|
||
"items": [
|
||
{"content": "Alice joined.", "strategy": "fast"},
|
||
{"content": "Bob left.", "strategy": "detailed"},
|
||
{"content": "Carol arrived."}, # no strategy — falls back to bank default
|
||
],
|
||
}
|
||
)
|
||
|
||
assert request.items[0].strategy == "fast"
|
||
assert request.items[1].strategy == "detailed"
|
||
assert request.items[2].strategy is None
|
||
|
||
# Simulate grouping logic from api_retain handler
|
||
strategy_groups: dict = {}
|
||
for item in request.items:
|
||
strategy_groups.setdefault(item.strategy, []).append(item.content)
|
||
|
||
assert set(strategy_groups.keys()) == {"fast", "detailed", None}
|
||
assert strategy_groups["fast"] == ["Alice joined."]
|
||
assert strategy_groups["detailed"] == ["Bob left."]
|
||
assert strategy_groups[None] == ["Carol arrived."]
|
||
logger.info("✓ per-item strategy grouping works correctly")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_named_strategy_applied_end_to_end(memory, request_context):
|
||
"""
|
||
Integration test: a named strategy stored in bank config is actually applied
|
||
during retain_batch_async.
|
||
|
||
Regression test for the bug where strategy was passed through the HTTP layer
|
||
but the extraction mode override was silently ignored, always using the bank
|
||
default (e.g. 'concise') instead of the strategy's override (e.g. 'chunks').
|
||
"""
|
||
from hindsight_api.config_resolver import ConfigResolver
|
||
|
||
bank_id = f"test_strategy_e2e_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Seed the bank so the row exists before we write config to it
|
||
# (update_bank_config is a plain UPDATE — it silently no-ops on missing rows)
|
||
await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[{"content": "seed"}],
|
||
request_context=request_context,
|
||
)
|
||
|
||
# Now configure the bank with a named strategy that overrides to chunks
|
||
await memory._config_resolver.update_bank_config(
|
||
bank_id,
|
||
{
|
||
"retain_extraction_mode": "concise", # bank default
|
||
"retain_strategies": {
|
||
"chunks": {"retain_extraction_mode": "chunks"},
|
||
},
|
||
},
|
||
request_context,
|
||
)
|
||
|
||
contents = [{"content": "Alice deployed the new API on Monday."}]
|
||
|
||
# Retain using the named strategy
|
||
unit_ids_by_content, usage = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=contents,
|
||
strategy="chunks",
|
||
request_context=request_context,
|
||
return_usage=True,
|
||
)
|
||
|
||
# chunks produces exactly one fact per chunk (verbatim) and calls no LLM
|
||
assert usage.total_tokens == 0, f"chunks should use zero LLM tokens, got {usage.total_tokens}"
|
||
assert len(unit_ids_by_content) == 1
|
||
assert len(unit_ids_by_content[0]) == 1, "chunks should produce exactly one fact per content item"
|
||
|
||
# Verify the stored fact is the verbatim content
|
||
facts = await memory.recall_async(bank_id, "Alice", request_context=request_context)
|
||
assert any("Alice" in f.text for f in facts.results), "Verbatim content should be retrievable"
|
||
|
||
logger.info("✓ named strategy 'chunks' with chunks applied end-to-end: no LLM, verbatim storage")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_semantic_ann_uses_hnsw_index(memory, request_context):
|
||
"""
|
||
Test that Phase 1 ANN semantic search creates links between similar world
|
||
facts across batches. This exercises the per-fact_type partial HNSW index
|
||
and the placeholder-ID remap logic.
|
||
"""
|
||
bank_id = f"test_sem_ann_hnsw_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
# Batch 1: world facts about machine learning
|
||
unit_ids_1 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=(
|
||
"Deep learning models require large amounts of training data. "
|
||
"Gradient descent is the primary optimization algorithm used in neural networks."
|
||
),
|
||
context="ML knowledge base",
|
||
event_date=datetime(2024, 3, 1, tzinfo=timezone.utc),
|
||
request_context=request_context,
|
||
)
|
||
assert len(unit_ids_1) > 0, "Batch 1 should produce facts"
|
||
|
||
# Batch 2: similar ML world facts — Phase 1 ANN should link to batch 1
|
||
unit_ids_2 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content=(
|
||
"Neural networks learn by adjusting weights through backpropagation. "
|
||
"Training deep learning models requires GPUs for fast gradient computation."
|
||
),
|
||
context="ML knowledge base",
|
||
event_date=datetime(2024, 3, 2, tzinfo=timezone.utc),
|
||
request_context=request_context,
|
||
)
|
||
assert len(unit_ids_2) > 0, "Batch 2 should produce facts"
|
||
|
||
logger.info(f"Batch 1: {len(unit_ids_1)} facts, Batch 2: {len(unit_ids_2)} facts")
|
||
|
||
# Verify cross-batch semantic links exist
|
||
async with memory._pool.acquire() as conn:
|
||
cross_links = await conn.fetch(
|
||
"""
|
||
SELECT from_unit_id, to_unit_id, weight
|
||
FROM memory_links
|
||
WHERE from_unit_id::text = ANY($1)
|
||
AND link_type = 'semantic'
|
||
AND to_unit_id::text = ANY($2)
|
||
""",
|
||
unit_ids_2,
|
||
unit_ids_1,
|
||
)
|
||
|
||
logger.info(f"Cross-batch semantic links (batch2 -> batch1): {len(cross_links)}")
|
||
for link in cross_links:
|
||
logger.info(
|
||
f" {str(link['from_unit_id'])[:8]}... -> "
|
||
f"{str(link['to_unit_id'])[:8]}... (weight: {link['weight']:.3f})"
|
||
)
|
||
|
||
assert len(cross_links) > 0, (
|
||
"Phase 1 ANN should create semantic links between similar world facts "
|
||
"from different batches via the HNSW index with placeholder-ID remap."
|
||
)
|
||
|
||
# All weights must meet the similarity threshold
|
||
for link in cross_links:
|
||
assert link["weight"] >= 0.7, (
|
||
f"Semantic link weight {link['weight']:.3f} below threshold 0.7"
|
||
)
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_temporal_links_scoped_by_fact_type(memory, request_context):
|
||
"""
|
||
Test that temporal links only connect facts of the SAME fact_type.
|
||
|
||
World facts should not get temporal links to experience facts even when
|
||
their event dates fall within the time window.
|
||
"""
|
||
bank_id = f"test_temporal_scope_{datetime.now(timezone.utc).timestamp()}"
|
||
|
||
try:
|
||
base_date = datetime(2024, 5, 10, 12, 0, 0, tzinfo=timezone.utc)
|
||
|
||
# Store a world fact
|
||
world_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="Python 3.12 was released with significant performance improvements for the interpreter.",
|
||
context="tech news",
|
||
event_date=base_date,
|
||
fact_type_override="world",
|
||
request_context=request_context,
|
||
)
|
||
assert len(world_ids) > 0, "Should create world fact(s)"
|
||
|
||
# Store an experience fact at a nearby timestamp (same hour)
|
||
experience_ids = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="I upgraded all my projects to Python 3.12 and benchmarked the speed improvements.",
|
||
context="personal log",
|
||
event_date=base_date + timedelta(hours=1),
|
||
fact_type_override="experience",
|
||
request_context=request_context,
|
||
)
|
||
assert len(experience_ids) > 0, "Should create experience fact(s)"
|
||
|
||
# Store another world fact at a nearby timestamp so we can confirm
|
||
# same-type temporal links ARE created
|
||
world_ids_2 = await memory.retain_async(
|
||
bank_id=bank_id,
|
||
content="The Python Software Foundation announced long-term support plans for Python 3.12.",
|
||
context="tech news",
|
||
event_date=base_date + timedelta(hours=2),
|
||
fact_type_override="world",
|
||
request_context=request_context,
|
||
)
|
||
assert len(world_ids_2) > 0, "Should create second world fact(s)"
|
||
|
||
logger.info(
|
||
f"World1: {world_ids}, Experience: {experience_ids}, World2: {world_ids_2}"
|
||
)
|
||
|
||
async with memory._pool.acquire() as conn:
|
||
# Check that world facts DO have temporal links to each other
|
||
world_all = world_ids + world_ids_2
|
||
world_temporal = await conn.fetch(
|
||
"""
|
||
SELECT from_unit_id, to_unit_id, weight
|
||
FROM memory_links
|
||
WHERE from_unit_id::text = ANY($1)
|
||
AND to_unit_id::text = ANY($1)
|
||
AND link_type = 'temporal'
|
||
""",
|
||
world_all,
|
||
)
|
||
logger.info(f"World-to-world temporal links: {len(world_temporal)}")
|
||
assert len(world_temporal) > 0, (
|
||
"World facts with nearby dates should have temporal links to each other"
|
||
)
|
||
|
||
# Check that world facts do NOT have temporal links to experience facts
|
||
cross_type_links = await conn.fetch(
|
||
"""
|
||
SELECT from_unit_id, to_unit_id, weight
|
||
FROM memory_links
|
||
WHERE (
|
||
(from_unit_id::text = ANY($1) AND to_unit_id::text = ANY($2))
|
||
OR
|
||
(from_unit_id::text = ANY($2) AND to_unit_id::text = ANY($1))
|
||
)
|
||
AND link_type = 'temporal'
|
||
""",
|
||
world_all,
|
||
experience_ids,
|
||
)
|
||
logger.info(f"Cross-type temporal links (world<->experience): {len(cross_type_links)}")
|
||
assert len(cross_type_links) == 0, (
|
||
f"Temporal links should NOT cross fact types, but found {len(cross_type_links)} "
|
||
f"world<->experience links"
|
||
)
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Streaming chunk batching tests
|
||
# ---------------------------------------------------------------------------
|
||
|
||
import json
|
||
import uuid
|
||
from unittest.mock import patch
|
||
|
||
import pytest_asyncio
|
||
|
||
from hindsight_api.engine.llm_wrapper import TokenUsage
|
||
from hindsight_api.engine.memory_engine import MemoryEngine
|
||
from hindsight_api.engine.task_backend import SyncTaskBackend
|
||
|
||
|
||
def _make_mock_llm_call():
|
||
"""Create a mock LLM call function that returns deterministic facts."""
|
||
|
||
async def mock_llm_call(*args, **kwargs):
|
||
from hindsight_api.engine.consolidation.consolidator import _ConsolidationBatchResponse
|
||
|
||
if kwargs.get("scope") == "consolidation":
|
||
return_usage = kwargs.get("return_usage", False)
|
||
if return_usage:
|
||
return _ConsolidationBatchResponse(), TokenUsage(input_tokens=0, output_tokens=0)
|
||
return _ConsolidationBatchResponse()
|
||
|
||
messages = kwargs.get("messages", args[0] if args else [])
|
||
user_msg = messages[-1]["content"] if messages else ""
|
||
|
||
# Extract sentences from the content to generate one fact per sentence
|
||
sentences = [s.strip() for s in user_msg.split(".") if s.strip() and len(s.strip()) > 10]
|
||
num_facts = max(1, min(len(sentences), 10))
|
||
|
||
facts = []
|
||
for i in range(num_facts):
|
||
sentence = sentences[i] if i < len(sentences) else f"Fact {i}"
|
||
facts.append({
|
||
"what": sentence[:200],
|
||
"when": "2024-06-15",
|
||
"where": "N/A",
|
||
"who": "N/A",
|
||
"why": "N/A",
|
||
"fact_type": "world",
|
||
"entities": [{"text": f"Entity{i}"}],
|
||
"causal_relations": [],
|
||
})
|
||
|
||
response_dict = {"facts": facts}
|
||
return_usage = kwargs.get("return_usage", False)
|
||
if return_usage:
|
||
usage = TokenUsage(
|
||
input_tokens=len(user_msg) // 4,
|
||
output_tokens=len(json.dumps(response_dict)) // 4,
|
||
)
|
||
return response_dict, usage
|
||
return response_dict
|
||
|
||
return mock_llm_call
|
||
|
||
|
||
@pytest_asyncio.fixture(scope="function")
|
||
async def memory_mock_llm(pg0_db_url, embeddings, cross_encoder, query_analyzer):
|
||
"""MemoryEngine with mock LLM for streaming tests."""
|
||
mem = MemoryEngine(
|
||
db_url=pg0_db_url,
|
||
memory_llm_provider="openai",
|
||
memory_llm_api_key="mock-key",
|
||
memory_llm_model="gpt-4",
|
||
embeddings=embeddings,
|
||
cross_encoder=cross_encoder,
|
||
query_analyzer=query_analyzer,
|
||
pool_min_size=1,
|
||
pool_max_size=5,
|
||
run_migrations=False,
|
||
skip_llm_verification=True,
|
||
task_backend=SyncTaskBackend(),
|
||
)
|
||
await mem.initialize()
|
||
yield mem
|
||
try:
|
||
if mem._pool and not mem._pool._closing:
|
||
await mem.close()
|
||
except Exception:
|
||
pass
|
||
|
||
|
||
def _generate_chunky_content(num_chunks: int, chunk_size: int = 3000) -> str:
|
||
"""Generate content that will produce approximately num_chunks chunks.
|
||
|
||
Each chunk is chunk_size characters, separated by double newlines.
|
||
"""
|
||
base_sentences = [
|
||
"Alice works as a senior engineer at TechCorp in San Francisco.",
|
||
"Bob joined the marketing team last month from Chicago.",
|
||
"The project deadline was extended to December 15th.",
|
||
"Sarah mentioned she is planning a trip to Tokyo next month.",
|
||
"The quarterly budget review showed a 15% increase in revenue.",
|
||
"Mike suggested exploring alternative cloud providers.",
|
||
"The client feedback from beta testing was positive overall.",
|
||
"Emily started learning Rust programming language last week.",
|
||
"The new office will be located in the financial district.",
|
||
"David presented the annual technology roadmap to stakeholders.",
|
||
]
|
||
|
||
chunks = []
|
||
for chunk_idx in range(num_chunks):
|
||
# Generate enough text for one chunk
|
||
lines = []
|
||
chars = 0
|
||
line_idx = 0
|
||
while chars < chunk_size - 100:
|
||
sentence = f"[Chunk {chunk_idx}, Line {line_idx}] {base_sentences[line_idx % len(base_sentences)]}"
|
||
lines.append(sentence)
|
||
chars += len(sentence) + 1
|
||
line_idx += 1
|
||
chunks.append("\n".join(lines))
|
||
|
||
return "\n\n".join(chunks)
|
||
|
||
|
||
def _set_chunk_batch_size(memory: MemoryEngine, batch_size: int) -> None:
|
||
"""Set retain_chunk_batch_size on the config resolver's global config."""
|
||
memory._config_resolver._global_config.retain_chunk_batch_size = batch_size
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_streaming_chunk_batching_produces_same_facts(memory_mock_llm, request_context):
|
||
"""
|
||
Retain a medium document (~10 chunks) with batch_size=3.
|
||
Verify all facts are extracted (streaming should not lose facts).
|
||
"""
|
||
memory = memory_mock_llm
|
||
_set_chunk_batch_size(memory, 3)
|
||
bank_id = f"test_streaming_{uuid.uuid4().hex[:8]}"
|
||
document_id = f"streaming_doc_{uuid.uuid4().hex[:8]}"
|
||
|
||
# Generate content that produces ~10 chunks at default chunk_size (3000 chars)
|
||
content = _generate_chunky_content(num_chunks=10, chunk_size=3000)
|
||
|
||
mock_llm_call = _make_mock_llm_call()
|
||
|
||
try:
|
||
with patch("hindsight_api.engine.llm_wrapper.LLMProvider.call", new=mock_llm_call):
|
||
# Retain with streaming enabled (batch_size=3, so 10 chunks -> 4 mini-batches)
|
||
result = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[{
|
||
"content": content,
|
||
"context": "streaming test",
|
||
"event_date": datetime(2024, 6, 15, tzinfo=timezone.utc),
|
||
}],
|
||
document_id=document_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
streaming_unit_ids = result[0] if result else []
|
||
logger.info(f"Streaming produced {len(streaming_unit_ids)} facts")
|
||
assert len(streaming_unit_ids) > 0, "Streaming should produce facts"
|
||
|
||
# Verify facts are in the DB
|
||
async with memory._pool.acquire() as conn:
|
||
fact_count = await conn.fetchval(
|
||
"SELECT COUNT(*) FROM memory_units WHERE bank_id = $1",
|
||
bank_id,
|
||
)
|
||
assert fact_count == len(streaming_unit_ids), (
|
||
f"DB has {fact_count} facts, but streaming returned {len(streaming_unit_ids)} unit_ids"
|
||
)
|
||
|
||
# Verify the document was tracked
|
||
doc = await conn.fetchrow(
|
||
"SELECT id FROM documents WHERE bank_id = $1 AND id = $2",
|
||
bank_id, document_id,
|
||
)
|
||
assert doc is not None, "Document should be tracked in DB"
|
||
|
||
# Verify chunks were stored with correct indices
|
||
chunk_count = await conn.fetchval(
|
||
"SELECT COUNT(*) FROM chunks WHERE bank_id = $1 AND document_id = $2",
|
||
bank_id, document_id,
|
||
)
|
||
assert chunk_count > 0, "Chunks should be stored in DB"
|
||
logger.info(f"Stored {chunk_count} chunks for document {document_id}")
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_streaming_chunk_batching_recovery(memory_mock_llm, request_context):
|
||
"""
|
||
Test recovery: retain a document with streaming, then retain the same
|
||
document again. Delta retain should detect existing chunks and skip
|
||
re-extraction. Fact count should be unchanged (no duplicates).
|
||
"""
|
||
memory = memory_mock_llm
|
||
_set_chunk_batch_size(memory, 3)
|
||
bank_id = f"test_streaming_recovery_{uuid.uuid4().hex[:8]}"
|
||
document_id = f"recovery_doc_{uuid.uuid4().hex[:8]}"
|
||
|
||
content = _generate_chunky_content(num_chunks=9, chunk_size=3000)
|
||
|
||
mock_llm_call = _make_mock_llm_call()
|
||
|
||
try:
|
||
# First retain — streaming mode
|
||
with patch("hindsight_api.engine.llm_wrapper.LLMProvider.call", new=mock_llm_call):
|
||
result1 = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[{
|
||
"content": content,
|
||
"context": "recovery test",
|
||
"event_date": datetime(2024, 6, 15, tzinfo=timezone.utc),
|
||
}],
|
||
document_id=document_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
first_unit_ids = result1[0] if result1 else []
|
||
assert len(first_unit_ids) > 0, "First retain should produce facts"
|
||
|
||
async with memory._pool.acquire() as conn:
|
||
first_fact_count = await conn.fetchval(
|
||
"SELECT COUNT(*) FROM memory_units WHERE bank_id = $1",
|
||
bank_id,
|
||
)
|
||
|
||
logger.info(f"First retain: {first_fact_count} facts")
|
||
|
||
# Second retain — same document, same content (should be a no-op via delta retain)
|
||
with patch("hindsight_api.engine.llm_wrapper.LLMProvider.call", new=mock_llm_call):
|
||
result2 = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[{
|
||
"content": content,
|
||
"context": "recovery test",
|
||
"event_date": datetime(2024, 6, 15, tzinfo=timezone.utc),
|
||
}],
|
||
document_id=document_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
async with memory._pool.acquire() as conn:
|
||
second_fact_count = await conn.fetchval(
|
||
"SELECT COUNT(*) FROM memory_units WHERE bank_id = $1",
|
||
bank_id,
|
||
)
|
||
|
||
logger.info(f"Second retain: {second_fact_count} facts")
|
||
|
||
# Fact count should be the same (delta retain skipped unchanged chunks)
|
||
assert second_fact_count == first_fact_count, (
|
||
f"Second retain should not create duplicates: first={first_fact_count}, second={second_fact_count}"
|
||
)
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_streaming_disabled_for_small_docs(memory_mock_llm, request_context):
|
||
"""
|
||
Retain a small document (2 chunks) with batch_size=500.
|
||
Verify it uses the non-streaming path (no batching overhead).
|
||
"""
|
||
memory = memory_mock_llm
|
||
_set_chunk_batch_size(memory, 500)
|
||
bank_id = f"test_streaming_small_{uuid.uuid4().hex[:8]}"
|
||
document_id = f"small_doc_{uuid.uuid4().hex[:8]}"
|
||
|
||
# Generate content that produces ~2 chunks
|
||
content = _generate_chunky_content(num_chunks=2, chunk_size=3000)
|
||
|
||
mock_llm_call = _make_mock_llm_call()
|
||
|
||
try:
|
||
with patch("hindsight_api.engine.llm_wrapper.LLMProvider.call", new=mock_llm_call):
|
||
# batch_size=500 >> 2 chunks, so non-streaming path should be used
|
||
result = await memory.retain_batch_async(
|
||
bank_id=bank_id,
|
||
contents=[{
|
||
"content": content,
|
||
"context": "small doc test",
|
||
"event_date": datetime(2024, 6, 15, tzinfo=timezone.utc),
|
||
}],
|
||
document_id=document_id,
|
||
request_context=request_context,
|
||
)
|
||
|
||
unit_ids = result[0] if result else []
|
||
logger.info(f"Small doc produced {len(unit_ids)} facts")
|
||
assert len(unit_ids) > 0, "Should produce facts even through non-streaming path"
|
||
|
||
# Verify the document was tracked
|
||
async with memory._pool.acquire() as conn:
|
||
doc = await conn.fetchrow(
|
||
"SELECT id FROM documents WHERE bank_id = $1 AND id = $2",
|
||
bank_id, document_id,
|
||
)
|
||
assert doc is not None, "Document should be tracked"
|
||
|
||
finally:
|
||
await memory.delete_bank(bank_id, request_context=request_context)
|