This commit is contained in:
Nicolò Boschi 2025-11-07 13:22:36 +01:00
parent 27bb335d7c
commit 7a9f4ee33c
2 changed files with 380 additions and 244 deletions

View file

@ -154,6 +154,11 @@ class TemporalSemanticMemory(
# we use ~20-40 connections max, staying well within pool limits
self._search_semaphore = asyncio.Semaphore(10)
# Backpressure for put operations: limit concurrent puts to prevent database contention
# Each put_batch holds a connection for the entire transaction, so we limit to 5
# concurrent puts to avoid connection pool exhaustion and reduce write contention
self._put_semaphore = asyncio.Semaphore(5)
async def _handle_access_count_update(self, task_dict: Dict[str, Any]):
"""
Handler for access count update tasks.
@ -177,6 +182,36 @@ class TemporalSemanticMemory(
except Exception as e:
logger.error(f"Access count handler: Error updating access counts: {e}")
async def _handle_batch_put(self, task_dict: Dict[str, Any]):
"""
Handler for batch put tasks.
Args:
task_dict: Dict with 'agent_id', 'contents', 'document_id', 'document_metadata', 'upsert'
"""
try:
agent_id = task_dict.get('agent_id')
contents = task_dict.get('contents', [])
document_id = task_dict.get('document_id')
document_metadata = task_dict.get('document_metadata')
upsert = task_dict.get('upsert', False)
logger.info(f"[BATCH_PUT_TASK] Starting background batch put for agent_id={agent_id}, {len(contents)} items")
await self.put_batch_async(
agent_id=agent_id,
contents=contents,
document_id=document_id,
document_metadata=document_metadata,
upsert=upsert
)
logger.info(f"[BATCH_PUT_TASK] Completed background batch put for agent_id={agent_id}")
except Exception as e:
logger.error(f"Batch put handler: Error processing batch put: {e}")
import traceback
traceback.print_exc()
async def execute_task(self, task_dict: Dict[str, Any]):
"""
Execute a task by routing it to the appropriate handler.
@ -196,6 +231,8 @@ class TemporalSemanticMemory(
await self._handle_reinforce_opinion(task_dict)
elif task_type == 'form_opinion':
await self._handle_form_opinion(task_dict)
elif task_type == 'batch_put':
await self._handle_batch_put(task_dict)
else:
logger.error(f"Unknown task type: {task_type}")
@ -591,286 +628,290 @@ class TemporalSemanticMemory(
Assumes contents are already appropriately sized (< 50k chars).
Called by put_batch_async after chunking large batches.
Uses semaphore for backpressure to limit concurrent puts.
"""
start_time = time.time()
total_chars = sum(len(item.get("content", "")) for item in contents)
# Backpressure: limit concurrent puts to prevent database contention
async with self._put_semaphore:
start_time = time.time()
total_chars = sum(len(item.get("content", "")) for item in contents)
# Buffer all logs to avoid interleaving
log_buffer = []
log_buffer.append(f"{'='*60}")
log_buffer.append(f"PUT_BATCH_ASYNC START: {agent_id}")
log_buffer.append(f"Batch size: {len(contents)} content items, {total_chars:,} chars")
log_buffer.append(f"{'='*60}")
# Buffer all logs to avoid interleaving
log_buffer = []
log_buffer.append(f"{'='*60}")
log_buffer.append(f"PUT_BATCH_ASYNC START: {agent_id}")
log_buffer.append(f"Batch size: {len(contents)} content items, {total_chars:,} chars")
log_buffer.append(f"{'='*60}")
# Step 1: Extract facts from ALL contents in parallel
step_start = time.time()
# Step 1: Extract facts from ALL contents in parallel
step_start = time.time()
# Create tasks for parallel fact extraction using configured LLM
fact_extraction_tasks = []
for item in contents:
content = item["content"]
context = item.get("context", "")
event_date = item.get("event_date") or utcnow()
# Create tasks for parallel fact extraction using configured LLM
fact_extraction_tasks = []
for item in contents:
content = item["content"]
context = item.get("context", "")
event_date = item.get("event_date") or utcnow()
task = extract_facts(content, event_date, context, llm_config=self._llm_config)
fact_extraction_tasks.append((task, event_date, context))
task = extract_facts(content, event_date, context, llm_config=self._llm_config)
fact_extraction_tasks.append((task, event_date, context))
# Wait for all fact extractions to complete
all_fact_results = await asyncio.gather(*[task for task, _, _ in fact_extraction_tasks])
log_buffer.append(f"[1] Extract facts (parallel): {len(fact_extraction_tasks)} contents in {time.time() - step_start:.3f}s")
# Wait for all fact extractions to complete
all_fact_results = await asyncio.gather(*[task for task, _, _ in fact_extraction_tasks])
log_buffer.append(f"[1] Extract facts (parallel): {len(fact_extraction_tasks)} contents in {time.time() - step_start:.3f}s")
# Flatten and track which facts belong to which content
all_fact_texts = []
all_fact_dates = []
all_contexts = []
all_fact_entities = [] # NEW: Store LLM-extracted entities per fact
all_fact_types = [] # Store fact type (world or agent)
content_boundaries = [] # [(start_idx, end_idx), ...]
# Flatten and track which facts belong to which content
all_fact_texts = []
all_fact_dates = []
all_contexts = []
all_fact_entities = [] # NEW: Store LLM-extracted entities per fact
all_fact_types = [] # Store fact type (world or agent)
content_boundaries = [] # [(start_idx, end_idx), ...]
current_idx = 0
for i, ((_, event_date, context), fact_dicts) in enumerate(zip(fact_extraction_tasks, all_fact_results)):
start_idx = current_idx
current_idx = 0
for i, ((_, event_date, context), fact_dicts) in enumerate(zip(fact_extraction_tasks, all_fact_results)):
start_idx = current_idx
for fact_dict in fact_dicts:
all_fact_texts.append(fact_dict['fact'])
try:
from dateutil import parser as date_parser
fact_date = date_parser.isoparse(fact_dict['date'])
all_fact_dates.append(fact_date)
except Exception:
all_fact_dates.append(event_date)
all_contexts.append(context)
# Extract entities from fact (default to empty list if not present)
all_fact_entities.append(fact_dict.get('entities', []))
# Extract fact type (use override if provided, else use extracted type or default to 'world')
if fact_type_override:
all_fact_types.append(fact_type_override)
else:
all_fact_types.append(fact_dict.get('fact_type', 'world'))
for fact_dict in fact_dicts:
all_fact_texts.append(fact_dict['fact'])
try:
from dateutil import parser as date_parser
fact_date = date_parser.isoparse(fact_dict['date'])
all_fact_dates.append(fact_date)
except Exception:
all_fact_dates.append(event_date)
all_contexts.append(context)
# Extract entities from fact (default to empty list if not present)
all_fact_entities.append(fact_dict.get('entities', []))
# Extract fact type (use override if provided, else use extracted type or default to 'world')
if fact_type_override:
all_fact_types.append(fact_type_override)
else:
all_fact_types.append(fact_dict.get('fact_type', 'world'))
end_idx = current_idx + len(fact_dicts)
content_boundaries.append((start_idx, end_idx))
current_idx = end_idx
end_idx = current_idx + len(fact_dicts)
content_boundaries.append((start_idx, end_idx))
current_idx = end_idx
total_facts = len(all_fact_texts)
total_facts = len(all_fact_texts)
if total_facts == 0:
return [[] for _ in contents]
if total_facts == 0:
return [[] for _ in contents]
# Step 2: Augment fact texts with readable dates for better temporal matching
# This allows queries like "camping in June" to match facts that happened in June
augmented_texts = []
for fact_text, fact_date in zip(all_fact_texts, all_fact_dates):
# Format date in readable form
readable_date = self._format_readable_date(fact_date)
# Augment text with date for embedding (but store original text in DB)
augmented_text = f"{fact_text} (happened in {readable_date})"
augmented_texts.append(augmented_text)
# Step 2: Augment fact texts with readable dates for better temporal matching
# This allows queries like "camping in June" to match facts that happened in June
augmented_texts = []
for fact_text, fact_date in zip(all_fact_texts, all_fact_dates):
# Format date in readable form
readable_date = self._format_readable_date(fact_date)
# Augment text with date for embedding (but store original text in DB)
augmented_text = f"{fact_text} (happened in {readable_date})"
augmented_texts.append(augmented_text)
# Step 2b: Generate ALL embeddings in ONE batch using augmented texts (HUGE speedup!)
step_start = time.time()
all_embeddings = await self._generate_embeddings_batch(augmented_texts)
log_buffer.append(f"[2] Generate embeddings (parallel): {len(all_embeddings)} embeddings in {time.time() - step_start:.3f}s")
# Step 2b: Generate ALL embeddings in ONE batch using augmented texts (HUGE speedup!)
step_start = time.time()
all_embeddings = await self._generate_embeddings_batch(augmented_texts)
log_buffer.append(f"[2] Generate embeddings (parallel): {len(all_embeddings)} embeddings in {time.time() - step_start:.3f}s")
# Step 3: Process everything in ONE database transaction
logger.debug("Getting connection pool")
pool = await self._get_pool()
logger.debug("Acquiring connection from pool")
async with pool.acquire() as conn:
logger.debug("Starting transaction")
async with conn.transaction():
logger.debug("Inside transaction")
try:
# Handle document tracking and upsert
if document_id:
logger.debug(f"Handling document tracking for {document_id}")
import hashlib
import json
# Step 3: Process everything in ONE database transaction
logger.debug("Getting connection pool")
pool = await self._get_pool()
logger.debug("Acquiring connection from pool")
async with pool.acquire() as conn:
logger.debug("Starting transaction")
async with conn.transaction():
logger.debug("Inside transaction")
try:
# Handle document tracking and upsert
if document_id:
logger.debug(f"Handling document tracking for {document_id}")
import hashlib
import json
# Calculate content hash from all content items
combined_content = "\n".join([c.get("content", "") for c in contents])
content_hash = hashlib.sha256(combined_content.encode()).hexdigest()
# Calculate content hash from all content items
combined_content = "\n".join([c.get("content", "") for c in contents])
content_hash = hashlib.sha256(combined_content.encode()).hexdigest()
# If upsert, delete old document first (cascades to units and links)
if upsert:
deleted = await conn.fetchval(
"DELETE FROM documents WHERE id = $1 AND agent_id = $2 RETURNING id",
document_id, agent_id
# If upsert, delete old document first (cascades to units and links)
if upsert:
deleted = await conn.fetchval(
"DELETE FROM documents WHERE id = $1 AND agent_id = $2 RETURNING id",
document_id, agent_id
)
if deleted:
logger.debug(f"[3.1] Upsert: Deleted existing document '{document_id}' and all its units")
# Insert or update document
# Always use ON CONFLICT for idempotent behavior
await conn.execute(
"""
INSERT INTO documents (id, agent_id, original_text, content_hash, metadata)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (id, agent_id) DO UPDATE
SET original_text = EXCLUDED.original_text,
content_hash = EXCLUDED.content_hash,
metadata = EXCLUDED.metadata,
updated_at = NOW()
""",
document_id,
agent_id,
combined_content,
content_hash,
json.dumps(document_metadata or {})
)
if deleted:
logger.debug(f"[3.1] Upsert: Deleted existing document '{document_id}' and all its units")
logger.debug(f"[3.2] Document '{document_id}' stored/updated")
# Insert or update document
# Always use ON CONFLICT for idempotent behavior
await conn.execute(
# Deduplication check for all facts (batched by time window)
logger.debug("Starting deduplication check")
step_start = time.time()
# Group facts by event_date (rounded to 12-hour buckets) for batching
from collections import defaultdict
time_buckets = defaultdict(list)
for idx, (sentence, embedding, fact_date) in enumerate(zip(all_fact_texts, all_embeddings, all_fact_dates)):
# Round to 12-hour bucket to group similar times
bucket_key = fact_date.replace(hour=(fact_date.hour // 12) * 12, minute=0, second=0, microsecond=0)
time_buckets[bucket_key].append((idx, sentence, embedding, fact_date))
# Process each bucket in batch
all_is_duplicate = [False] * total_facts # Initialize all as not duplicate
for bucket_date, bucket_items in time_buckets.items():
indices = [item[0] for item in bucket_items]
sentences = [item[1] for item in bucket_items]
embeddings = [item[2] for item in bucket_items]
# Use bucket_date as representative for time window
dup_flags = await self._find_duplicate_facts_batch(
conn, agent_id, sentences, embeddings, bucket_date, time_window_hours=24
)
# Map results back to original indices
for idx, is_dup in zip(indices, dup_flags):
all_is_duplicate[idx] = is_dup
duplicates_filtered = sum(all_is_duplicate)
new_facts = total_facts - duplicates_filtered
logger.debug(f"Deduplication complete: {duplicates_filtered} duplicates filtered, {new_facts} new facts ({len(time_buckets)} time buckets)")
log_buffer.append(f"[3] Deduplication check: {duplicates_filtered} duplicates filtered, {new_facts} new facts in {time.time() - step_start:.3f}s")
# Filter out duplicates
filtered_sentences = [s for s, is_dup in zip(all_fact_texts, all_is_duplicate) if not is_dup]
filtered_embeddings = [e for e, is_dup in zip(all_embeddings, all_is_duplicate) if not is_dup]
filtered_dates = [d for d, is_dup in zip(all_fact_dates, all_is_duplicate) if not is_dup]
filtered_contexts = [c for c, is_dup in zip(all_contexts, all_is_duplicate) if not is_dup]
filtered_entities = [ents for ents, is_dup in zip(all_fact_entities, all_is_duplicate) if not is_dup]
filtered_fact_types = [ft for ft, is_dup in zip(all_fact_types, all_is_duplicate) if not is_dup]
if not filtered_sentences:
logger.debug(f"[PUT_BATCH_ASYNC] All facts were duplicates, returning empty")
return [[] for _ in contents]
# Batch insert ALL units
step_start = time.time()
# Convert embeddings to strings for asyncpg vector type
filtered_embeddings_str = [str(emb) for emb in filtered_embeddings]
# Prepare confidence scores (only for opinions)
# If fact_type is 'opinion' and no confidence_score provided, use default of 1.0
confidence_scores = [
confidence_score if confidence_score is not None else 1.0
if ft == 'opinion'
else None
for ft in filtered_fact_types
]
results = await conn.fetch(
"""
INSERT INTO documents (id, agent_id, original_text, content_hash, metadata)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (id, agent_id) DO UPDATE
SET original_text = EXCLUDED.original_text,
content_hash = EXCLUDED.content_hash,
metadata = EXCLUDED.metadata,
updated_at = NOW()
INSERT INTO memory_units (agent_id, document_id, text, context, embedding, event_date, fact_type, confidence_score, access_count)
SELECT * FROM unnest($1::text[], $2::text[], $3::text[], $4::text[], $5::vector[], $6::timestamptz[], $7::text[], $8::float[], $9::integer[])
RETURNING id
""",
document_id,
agent_id,
combined_content,
content_hash,
json.dumps(document_metadata or {})
[agent_id] * len(filtered_sentences),
[document_id] * len(filtered_sentences) if document_id else [None] * len(filtered_sentences),
filtered_sentences,
filtered_contexts,
filtered_embeddings_str,
filtered_dates,
filtered_fact_types,
confidence_scores,
[0] * len(filtered_sentences)
)
logger.debug(f"[3.2] Document '{document_id}' stored/updated")
# Deduplication check for all facts (batched by time window)
logger.debug("Starting deduplication check")
step_start = time.time()
created_unit_ids = [str(row['id']) for row in results]
logger.debug(f"Batch insert complete: {len(created_unit_ids)} units created")
log_buffer.append(f"[5] Batch insert units: {len(created_unit_ids)} units in {time.time() - step_start:.3f}s")
# Group facts by event_date (rounded to 12-hour buckets) for batching
from collections import defaultdict
time_buckets = defaultdict(list)
for idx, (sentence, embedding, fact_date) in enumerate(zip(all_fact_texts, all_embeddings, all_fact_dates)):
# Round to 12-hour bucket to group similar times
bucket_key = fact_date.replace(hour=(fact_date.hour // 12) * 12, minute=0, second=0, microsecond=0)
time_buckets[bucket_key].append((idx, sentence, embedding, fact_date))
# Process each bucket in batch
all_is_duplicate = [False] * total_facts # Initialize all as not duplicate
for bucket_date, bucket_items in time_buckets.items():
indices = [item[0] for item in bucket_items]
sentences = [item[1] for item in bucket_items]
embeddings = [item[2] for item in bucket_items]
# Use bucket_date as representative for time window
dup_flags = await self._find_duplicate_facts_batch(
conn, agent_id, sentences, embeddings, bucket_date, time_window_hours=24
# Process entities for ALL units
logger.debug("Processing entities")
step_start = time.time()
all_entity_links = await self._extract_entities_batch_optimized(
conn, agent_id, created_unit_ids, filtered_sentences, "", filtered_dates, filtered_entities, log_buffer
)
# Map results back to original indices
for idx, is_dup in zip(indices, dup_flags):
all_is_duplicate[idx] = is_dup
logger.debug(f"Entity processing complete: {len(all_entity_links)} links")
log_buffer.append(f"[6] Process entities (batched): {time.time() - step_start:.3f}s")
duplicates_filtered = sum(all_is_duplicate)
new_facts = total_facts - duplicates_filtered
logger.debug(f"Deduplication complete: {duplicates_filtered} duplicates filtered, {new_facts} new facts ({len(time_buckets)} time buckets)")
log_buffer.append(f"[3] Deduplication check: {duplicates_filtered} duplicates filtered, {new_facts} new facts in {time.time() - step_start:.3f}s")
# Create temporal links
logger.debug("Creating temporal links")
step_start = time.time()
await self._create_temporal_links_batch_per_fact(conn, agent_id, created_unit_ids, log_buffer=log_buffer)
logger.debug("Temporal links complete")
log_buffer.append(f"[7] Batch create temporal links: {time.time() - step_start:.3f}s")
# Filter out duplicates
filtered_sentences = [s for s, is_dup in zip(all_fact_texts, all_is_duplicate) if not is_dup]
filtered_embeddings = [e for e, is_dup in zip(all_embeddings, all_is_duplicate) if not is_dup]
filtered_dates = [d for d, is_dup in zip(all_fact_dates, all_is_duplicate) if not is_dup]
filtered_contexts = [c for c, is_dup in zip(all_contexts, all_is_duplicate) if not is_dup]
filtered_entities = [ents for ents, is_dup in zip(all_fact_entities, all_is_duplicate) if not is_dup]
filtered_fact_types = [ft for ft, is_dup in zip(all_fact_types, all_is_duplicate) if not is_dup]
# Create semantic links
logger.debug("Creating semantic links")
step_start = time.time()
await self._create_semantic_links_batch(conn, agent_id, created_unit_ids, filtered_embeddings, log_buffer=log_buffer)
logger.debug("Semantic links complete")
log_buffer.append(f"[8] Batch create semantic links: {time.time() - step_start:.3f}s")
if not filtered_sentences:
logger.debug(f"[PUT_BATCH_ASYNC] All facts were duplicates, returning empty")
return [[] for _ in contents]
# Insert entity links
logger.debug("Inserting entity links")
step_start = time.time()
if all_entity_links:
await self._insert_entity_links_batch(conn, all_entity_links)
logger.debug("Entity links inserted")
log_buffer.append(f"[9] Batch insert entity links: {time.time() - step_start:.3f}s")
# Batch insert ALL units
step_start = time.time()
# Convert embeddings to strings for asyncpg vector type
filtered_embeddings_str = [str(emb) for emb in filtered_embeddings]
# Prepare confidence scores (only for opinions)
# If fact_type is 'opinion' and no confidence_score provided, use default of 1.0
confidence_scores = [
confidence_score if confidence_score is not None else 1.0
if ft == 'opinion'
else None
for ft in filtered_fact_types
]
results = await conn.fetch(
"""
INSERT INTO memory_units (agent_id, document_id, text, context, embedding, event_date, fact_type, confidence_score, access_count)
SELECT * FROM unnest($1::text[], $2::text[], $3::text[], $4::text[], $5::vector[], $6::timestamptz[], $7::text[], $8::float[], $9::integer[])
RETURNING id
""",
[agent_id] * len(filtered_sentences),
[document_id] * len(filtered_sentences) if document_id else [None] * len(filtered_sentences),
filtered_sentences,
filtered_contexts,
filtered_embeddings_str,
filtered_dates,
filtered_fact_types,
confidence_scores,
[0] * len(filtered_sentences)
)
# Transaction auto-commits on success
commit_start = time.time()
logger.debug(f"[10] Commit: {time.time() - commit_start:.3f}s")
created_unit_ids = [str(row['id']) for row in results]
logger.debug(f"Batch insert complete: {len(created_unit_ids)} units created")
log_buffer.append(f"[5] Batch insert units: {len(created_unit_ids)} units in {time.time() - step_start:.3f}s")
# Map created unit IDs back to original content items
# Account for duplicates when mapping back
result_unit_ids = []
filtered_idx = 0
# Process entities for ALL units
logger.debug("Processing entities")
step_start = time.time()
all_entity_links = await self._extract_entities_batch_optimized(
conn, agent_id, created_unit_ids, filtered_sentences, "", filtered_dates, filtered_entities, log_buffer
)
logger.debug(f"Entity processing complete: {len(all_entity_links)} links")
log_buffer.append(f"[6] Process entities (batched): {time.time() - step_start:.3f}s")
for start_idx, end_idx in content_boundaries:
content_unit_ids = []
for i in range(start_idx, end_idx):
if not all_is_duplicate[i]:
content_unit_ids.append(created_unit_ids[filtered_idx])
filtered_idx += 1
result_unit_ids.append(content_unit_ids)
# Create temporal links
logger.debug("Creating temporal links")
step_start = time.time()
await self._create_temporal_links_batch_per_fact(conn, agent_id, created_unit_ids, log_buffer=log_buffer)
logger.debug("Temporal links complete")
log_buffer.append(f"[7] Batch create temporal links: {time.time() - step_start:.3f}s")
total_time = time.time() - start_time
log_buffer.append(f"{'='*60}")
log_buffer.append(f"PUT_BATCH_ASYNC COMPLETE: {len(created_unit_ids)} units from {len(contents)} contents in {total_time:.3f}s")
log_buffer.append(f"{'='*60}")
# Create semantic links
logger.debug("Creating semantic links")
step_start = time.time()
await self._create_semantic_links_batch(conn, agent_id, created_unit_ids, filtered_embeddings, log_buffer=log_buffer)
logger.debug("Semantic links complete")
log_buffer.append(f"[8] Batch create semantic links: {time.time() - step_start:.3f}s")
# Flush all logs at once to avoid interleaving
logger.info("\n" + "\n".join(log_buffer) + "\n")
# Insert entity links
logger.debug("Inserting entity links")
step_start = time.time()
if all_entity_links:
await self._insert_entity_links_batch(conn, all_entity_links)
logger.debug("Entity links inserted")
log_buffer.append(f"[9] Batch insert entity links: {time.time() - step_start:.3f}s")
# Trigger opinion reinforcement in background (non-blocking)
# Only trigger if there are entities in the new units
if any(filtered_entities):
await self._task_backend.submit_task({
'type': 'reinforce_opinion',
'agent_id': agent_id,
'created_unit_ids': created_unit_ids,
'unit_texts': filtered_sentences,
'unit_entities': filtered_entities
})
logger.debug("[PUT_BATCH_ASYNC] Opinion reinforcement task queued in background")
# Transaction auto-commits on success
commit_start = time.time()
logger.debug(f"[10] Commit: {time.time() - commit_start:.3f}s")
return result_unit_ids
# Map created unit IDs back to original content items
# Account for duplicates when mapping back
result_unit_ids = []
filtered_idx = 0
for start_idx, end_idx in content_boundaries:
content_unit_ids = []
for i in range(start_idx, end_idx):
if not all_is_duplicate[i]:
content_unit_ids.append(created_unit_ids[filtered_idx])
filtered_idx += 1
result_unit_ids.append(content_unit_ids)
total_time = time.time() - start_time
log_buffer.append(f"{'='*60}")
log_buffer.append(f"PUT_BATCH_ASYNC COMPLETE: {len(created_unit_ids)} units from {len(contents)} contents in {total_time:.3f}s")
log_buffer.append(f"{'='*60}")
# Flush all logs at once to avoid interleaving
logger.info("\n" + "\n".join(log_buffer) + "\n")
# Trigger opinion reinforcement in background (non-blocking)
# Only trigger if there are entities in the new units
if any(filtered_entities):
await self._task_backend.submit_task({
'type': 'reinforce_opinion',
'agent_id': agent_id,
'created_unit_ids': created_unit_ids,
'unit_texts': filtered_sentences,
'unit_entities': filtered_entities
})
logger.debug("[PUT_BATCH_ASYNC] Opinion reinforcement task queued in background")
return result_unit_ids
except Exception as e:
# Transaction auto-rolls back on exception
import traceback
traceback.print_exc()
raise Exception(f"Failed to store batch memory: {str(e)}")
except Exception as e:
# Transaction auto-rolls back on exception
import traceback
traceback.print_exc()
raise Exception(f"Failed to store batch memory: {str(e)}")
def search(
self,

View file

@ -206,6 +206,28 @@ class BatchPutResponse(BaseModel):
}
class BatchPutAsyncResponse(BaseModel):
"""Response model for async batch put endpoint."""
success: bool
message: str
agent_id: str
document_id: Optional[str] = None
items_count: int
queued: bool
class Config:
json_schema_extra = {
"example": {
"success": True,
"message": "Batch put task queued for background processing",
"agent_id": "user123",
"document_id": "conversation_123",
"items_count": 2,
"queued": True
}
}
class ThinkRequest(BaseModel):
"""Request model for think endpoint."""
query: str
@ -637,6 +659,79 @@ def _register_routes(app: FastAPI):
raise HTTPException(status_code=500, detail=str(e))
@app.post(
"/api/memories/batch_async",
response_model=BatchPutAsyncResponse,
tags=["Memory Storage"],
summary="Store multiple memories asynchronously",
description="""
Store multiple memory items in batch asynchronously using the task backend.
This endpoint returns immediately after queuing the task, without waiting for completion.
The actual processing happens in the background.
Features:
- Immediate response (non-blocking)
- Background processing via task queue
- Efficient batch processing
- Automatic fact extraction from natural language
- Entity recognition and linking
- Document tracking with optional upsert
- Temporal and semantic linking
The system automatically:
1. Queues the batch put task
2. Returns immediately with success=True, queued=True
3. Processes in background: extracts facts, generates embeddings, creates links
"""
)
async def api_batch_put_async(request: BatchPutRequest):
try:
# Validate agent_id - prevent writing to reserved agents
RESERVED_AGENT_IDS = {"locomo"}
if request.agent_id in RESERVED_AGENT_IDS:
raise HTTPException(
status_code=403,
detail=f"Cannot write to reserved agent_id '{request.agent_id}'. Reserved agents: {', '.join(RESERVED_AGENT_IDS)}"
)
# Prepare contents for put_batch_async
contents = []
for item in request.items:
content_dict = {"content": item.content}
if item.event_date:
content_dict["event_date"] = item.event_date
if item.context:
content_dict["context"] = item.context
contents.append(content_dict)
# Submit task to background queue
await app.state.memory._task_backend.submit_task({
'type': 'batch_put',
'agent_id': request.agent_id,
'contents': contents,
'document_id': request.document_id,
'document_metadata': request.document_metadata,
'upsert': request.upsert
})
logging.info(f"Batch put task queued for agent_id={request.agent_id}, {len(contents)} items")
return BatchPutAsyncResponse(
success=True,
message=f"Batch put task queued for background processing ({len(contents)} items)",
agent_id=request.agent_id,
document_id=request.document_id,
items_count=len(contents),
queued=True
)
except Exception as e:
import traceback
error_detail = f"{str(e)}\n\nTraceback:\n{traceback.format_exc()}"
print(f"Error in /api/memories/batch_async: {error_detail}")
raise HTTPException(status_code=500, detail=str(e))
@app.delete(
"/api/memory/{unit_id}",
tags=["Memory Storage"],