Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
25ea225b8e | ||
|
|
21b880aa8d | ||
|
|
89c914cc0b |
@@ -285,6 +285,7 @@ ENV_FILE_DELETE_AFTER_RETAIN = "HINDSIGHT_API_FILE_DELETE_AFTER_RETAIN"
|
||||
# Observations settings (consolidated knowledge from facts)
|
||||
ENV_ENABLE_OBSERVATIONS = "HINDSIGHT_API_ENABLE_OBSERVATIONS"
|
||||
ENV_CONSOLIDATION_BATCH_SIZE = "HINDSIGHT_API_CONSOLIDATION_BATCH_SIZE"
|
||||
ENV_CONSOLIDATION_LLM_BATCH_SIZE = "HINDSIGHT_API_CONSOLIDATION_LLM_BATCH_SIZE"
|
||||
ENV_CONSOLIDATION_MAX_TOKENS = "HINDSIGHT_API_CONSOLIDATION_MAX_TOKENS"
|
||||
ENV_OBSERVATIONS_MISSION = "HINDSIGHT_API_OBSERVATIONS_MISSION"
|
||||
|
||||
@@ -426,7 +427,8 @@ DEFAULT_FILE_DELETE_AFTER_RETAIN = True # Delete file bytes after retain (saves
|
||||
# Observations defaults (consolidated knowledge from facts)
|
||||
DEFAULT_ENABLE_OBSERVATIONS = True # Observations enabled by default
|
||||
DEFAULT_CONSOLIDATION_BATCH_SIZE = 50 # Memories to load per batch (internal memory optimization)
|
||||
DEFAULT_CONSOLIDATION_MAX_TOKENS = 1024 # Max tokens for recall when finding related observations
|
||||
DEFAULT_CONSOLIDATION_LLM_BATCH_SIZE = 8 # Facts per LLM call (1 = no batching; >1 = batch mode)
|
||||
DEFAULT_CONSOLIDATION_MAX_TOKENS = 512 # Max tokens for recall when finding related observations
|
||||
DEFAULT_OBSERVATIONS_MISSION = None # Declarative spec of what observations are for this bank
|
||||
|
||||
# Database migrations
|
||||
@@ -679,6 +681,7 @@ class HindsightConfig:
|
||||
# Observations settings (consolidated knowledge from facts)
|
||||
enable_observations: bool
|
||||
consolidation_batch_size: int
|
||||
consolidation_llm_batch_size: int
|
||||
consolidation_max_tokens: int
|
||||
observations_mission: str | None
|
||||
|
||||
@@ -1097,6 +1100,9 @@ class HindsightConfig:
|
||||
consolidation_batch_size=int(
|
||||
os.getenv(ENV_CONSOLIDATION_BATCH_SIZE, str(DEFAULT_CONSOLIDATION_BATCH_SIZE))
|
||||
),
|
||||
consolidation_llm_batch_size=int(
|
||||
os.getenv(ENV_CONSOLIDATION_LLM_BATCH_SIZE, str(DEFAULT_CONSOLIDATION_LLM_BATCH_SIZE))
|
||||
),
|
||||
consolidation_max_tokens=int(
|
||||
os.getenv(ENV_CONSOLIDATION_MAX_TOKENS, str(DEFAULT_CONSOLIDATION_MAX_TOKENS))
|
||||
),
|
||||
|
||||
@@ -15,6 +15,7 @@ import json
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
@@ -23,7 +24,7 @@ from pydantic import BaseModel
|
||||
from ...config import get_config
|
||||
from ..memory_engine import fq_table
|
||||
from ..retain import embedding_utils
|
||||
from .prompts import build_consolidation_prompt
|
||||
from .prompts import build_batch_consolidation_prompt
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from asyncpg import Connection
|
||||
@@ -35,15 +36,34 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class _ConsolidationAction(BaseModel):
|
||||
action: str # "update" | "create"
|
||||
class _CreateAction(BaseModel):
|
||||
text: str
|
||||
reason: str = ""
|
||||
learning_id: str | None = None # required for "update" actions
|
||||
source_fact_ids: list[str] # memory UUIDs from the NEW FACTS list
|
||||
|
||||
|
||||
class _ConsolidationResponse(BaseModel):
|
||||
actions: list[_ConsolidationAction]
|
||||
class _UpdateAction(BaseModel):
|
||||
text: str
|
||||
observation_id: str # UUID of the existing observation to update
|
||||
source_fact_ids: list[str] # memory UUIDs from the NEW FACTS list
|
||||
|
||||
|
||||
class _DeleteAction(BaseModel):
|
||||
observation_id: str # UUID of the observation to remove
|
||||
|
||||
|
||||
class _ConsolidationBatchResponse(BaseModel):
|
||||
creates: list[_CreateAction] = []
|
||||
updates: list[_UpdateAction] = []
|
||||
deletes: list[_DeleteAction] = []
|
||||
|
||||
|
||||
@dataclass
|
||||
class _BatchLLMResult:
|
||||
creates: list[_CreateAction] = field(default_factory=list)
|
||||
updates: list[_UpdateAction] = field(default_factory=list)
|
||||
deletes: list[_DeleteAction] = field(default_factory=list)
|
||||
obs_count: int = 0
|
||||
prompt_chars: int = 0
|
||||
|
||||
|
||||
class ConsolidationPerfLog:
|
||||
@@ -54,6 +74,9 @@ class ConsolidationPerfLog:
|
||||
self.start_time = time.time()
|
||||
self.lines: list[str] = []
|
||||
self.timings: dict[str, float] = {}
|
||||
self.llm_calls: int = 0
|
||||
self.total_obs_in_context: int = 0
|
||||
self.total_prompt_chars: int = 0
|
||||
|
||||
def log(self, message: str) -> None:
|
||||
"""Add a log line."""
|
||||
@@ -66,6 +89,12 @@ class ConsolidationPerfLog:
|
||||
else:
|
||||
self.timings[key] = duration
|
||||
|
||||
def record_llm_call(self, obs_count: int, prompt_chars: int) -> None:
|
||||
"""Record stats for a single LLM call."""
|
||||
self.llm_calls += 1
|
||||
self.total_obs_in_context += obs_count
|
||||
self.total_prompt_chars += prompt_chars
|
||||
|
||||
def flush(self) -> None:
|
||||
"""Flush all log lines to the logger."""
|
||||
total_time = time.time() - self.start_time
|
||||
@@ -98,6 +127,7 @@ async def run_consolidation_job(
|
||||
config = await memory_engine._config_resolver.resolve_full_config(bank_id, request_context)
|
||||
perf = ConsolidationPerfLog(bank_id)
|
||||
max_memories_per_batch = config.consolidation_batch_size
|
||||
llm_batch_size = max(1, config.consolidation_llm_batch_size)
|
||||
|
||||
# Check if consolidation is enabled
|
||||
if not config.enable_observations:
|
||||
@@ -156,15 +186,8 @@ async def run_consolidation_job(
|
||||
# Track all unique tags from consolidated memories for mental model refresh filtering
|
||||
consolidated_tags: set[str] = set()
|
||||
|
||||
batch_num = 0
|
||||
last_progress_timings = {} # Track timings at last progress log
|
||||
llm_batch_num = 0
|
||||
while True:
|
||||
batch_num += 1
|
||||
batch_start = time.time()
|
||||
|
||||
# Snapshot timings at batch start for per-batch calculation
|
||||
batch_start_timings = perf.timings.copy()
|
||||
|
||||
# Fetch next batch of unconsolidated memories
|
||||
async with pool.acquire() as conn:
|
||||
t0 = time.time()
|
||||
@@ -186,95 +209,90 @@ async def run_consolidation_job(
|
||||
if not memories:
|
||||
break # No more unconsolidated memories
|
||||
|
||||
for memory in memories:
|
||||
mem_start = time.time()
|
||||
# Group memories by exact tag set before batching — security requirement:
|
||||
# memories with different tags must never share an LLM call.
|
||||
tag_groups: dict[tuple[str, ...], list[dict[str, Any]]] = {}
|
||||
for m in memories:
|
||||
tag_key = tuple(sorted(m.get("tags") or []))
|
||||
tag_groups.setdefault(tag_key, []).append(dict(m))
|
||||
|
||||
# Track tags from this memory for mental model refresh filtering
|
||||
memory_tags = memory.get("tags") or []
|
||||
if memory_tags:
|
||||
consolidated_tags.update(memory_tags)
|
||||
# Flatten into LLM batches respecting both tag groups and llm_batch_size
|
||||
llm_batches: list[list[dict[str, Any]]] = []
|
||||
for group in tag_groups.values():
|
||||
for i in range(0, len(group), llm_batch_size):
|
||||
llm_batches.append(group[i : i + llm_batch_size])
|
||||
|
||||
for llm_batch in llm_batches:
|
||||
llm_batch_num += 1
|
||||
llm_batch_start = time.time()
|
||||
|
||||
# Snapshot perf and stats before this LLM batch
|
||||
snap_timings = perf.timings.copy()
|
||||
snap_llm_calls = perf.llm_calls
|
||||
snap_total_chars = perf.total_prompt_chars
|
||||
snap_stats = stats.copy()
|
||||
|
||||
# Track tags for mental model refresh filtering
|
||||
for memory in llm_batch:
|
||||
memory_tags = memory.get("tags") or []
|
||||
if memory_tags:
|
||||
consolidated_tags.update(memory_tags)
|
||||
|
||||
# Process the memory (uses its own connection internally)
|
||||
async with pool.acquire() as conn:
|
||||
result = await _process_memory(
|
||||
results = await _process_memory_batch(
|
||||
conn=conn,
|
||||
memory_engine=memory_engine,
|
||||
bank_id=bank_id,
|
||||
memory=dict(memory),
|
||||
memories=llm_batch,
|
||||
request_context=request_context,
|
||||
perf=perf,
|
||||
config=config,
|
||||
)
|
||||
|
||||
# Mark memory as consolidated (committed immediately)
|
||||
await conn.execute(
|
||||
f"""
|
||||
UPDATE {fq_table("memory_units")}
|
||||
SET consolidated_at = NOW()
|
||||
WHERE id = $1
|
||||
""",
|
||||
memory["id"],
|
||||
await conn.executemany(
|
||||
f"UPDATE {fq_table('memory_units')} SET consolidated_at = NOW() WHERE id = $1",
|
||||
[(m["id"],) for m in llm_batch],
|
||||
)
|
||||
|
||||
mem_time = time.time() - mem_start
|
||||
perf.record_timing("process_memory_total", mem_time)
|
||||
for result in results:
|
||||
stats["memories_processed"] += 1
|
||||
action = result.get("action")
|
||||
if action == "created":
|
||||
stats["observations_created"] += 1
|
||||
stats["actions_executed"] += 1
|
||||
elif action == "updated":
|
||||
stats["observations_updated"] += 1
|
||||
stats["actions_executed"] += 1
|
||||
elif action == "merged":
|
||||
stats["observations_merged"] += 1
|
||||
stats["actions_executed"] += 1
|
||||
elif action == "multiple":
|
||||
stats["observations_created"] += result.get("created", 0)
|
||||
stats["observations_updated"] += result.get("updated", 0)
|
||||
stats["observations_merged"] += result.get("merged", 0)
|
||||
stats["actions_executed"] += result.get("total_actions", 0)
|
||||
elif action == "skipped":
|
||||
stats["skipped"] += 1
|
||||
|
||||
stats["memories_processed"] += 1
|
||||
|
||||
action = result.get("action")
|
||||
if action == "created":
|
||||
stats["observations_created"] += 1
|
||||
stats["actions_executed"] += 1
|
||||
elif action == "updated":
|
||||
stats["observations_updated"] += 1
|
||||
stats["actions_executed"] += 1
|
||||
elif action == "merged":
|
||||
stats["observations_merged"] += 1
|
||||
stats["actions_executed"] += 1
|
||||
elif action == "multiple":
|
||||
stats["observations_created"] += result.get("created", 0)
|
||||
stats["observations_updated"] += result.get("updated", 0)
|
||||
stats["observations_merged"] += result.get("merged", 0)
|
||||
stats["actions_executed"] += result.get("total_actions", 0)
|
||||
elif action == "skipped":
|
||||
stats["skipped"] += 1
|
||||
|
||||
# Log progress periodically with timing breakdown
|
||||
if stats["memories_processed"] % 10 == 0:
|
||||
# Calculate timing deltas since last progress log
|
||||
timing_parts = []
|
||||
for key in ["recall", "llm", "embedding", "db_write"]:
|
||||
if key in perf.timings:
|
||||
delta = perf.timings[key] - last_progress_timings.get(key, 0)
|
||||
timing_parts.append(f"{key}={delta:.2f}s")
|
||||
|
||||
timing_str = f" | {', '.join(timing_parts)}" if timing_parts else ""
|
||||
logger.info(
|
||||
f"[CONSOLIDATION] bank={bank_id} progress: "
|
||||
f"{stats['memories_processed']}/{total_count} memories processed{timing_str}"
|
||||
)
|
||||
|
||||
# Update last progress snapshot
|
||||
last_progress_timings = perf.timings.copy()
|
||||
|
||||
batch_time = time.time() - batch_start
|
||||
perf.log(
|
||||
f"[2] Batch {batch_num}: {len(memories)} memories in {batch_time:.3f}s "
|
||||
f"(avg {batch_time / len(memories):.3f}s/memory)"
|
||||
)
|
||||
|
||||
# Log timing breakdown after each batch (delta from batch start)
|
||||
timing_parts = []
|
||||
for key in ["recall", "llm", "embedding", "db_write"]:
|
||||
if key in perf.timings:
|
||||
delta = perf.timings[key] - batch_start_timings.get(key, 0)
|
||||
timing_parts.append(f"{key}={delta:.3f}s")
|
||||
|
||||
if timing_parts:
|
||||
avg_per_memory = batch_time / len(memories) if memories else 0
|
||||
# Per-LLM-batch log
|
||||
llm_batch_time = time.time() - llm_batch_start
|
||||
timing_parts = []
|
||||
for key in ["recall", "llm", "embedding", "db_write"]:
|
||||
if key in perf.timings:
|
||||
delta = perf.timings[key] - snap_timings.get(key, 0)
|
||||
timing_parts.append(f"{key}={delta:.3f}s")
|
||||
input_tokens = int((perf.total_prompt_chars - snap_total_chars) / 4)
|
||||
batch_created = stats["observations_created"] - snap_stats["observations_created"]
|
||||
batch_updated = stats["observations_updated"] - snap_stats["observations_updated"]
|
||||
batch_skipped = stats["skipped"] - snap_stats["skipped"]
|
||||
llm_calls_made = perf.llm_calls - snap_llm_calls
|
||||
logger.info(
|
||||
f"[CONSOLIDATION] bank={bank_id} batch {batch_num}/{len(memories)} memories: "
|
||||
f"{', '.join(timing_parts)} | avg={avg_per_memory:.3f}s/memory"
|
||||
f"[CONSOLIDATION] bank={bank_id} llm_batch #{llm_batch_num}"
|
||||
f" ({len(llm_batch)} memories, {llm_calls_made} llm calls)"
|
||||
f" | {stats['memories_processed']}/{total_count} processed"
|
||||
f" | {', '.join(timing_parts)}"
|
||||
f" | created={batch_created} updated={batch_updated} skipped={batch_skipped}"
|
||||
f" | input_tokens=~{input_tokens}"
|
||||
f" | avg={llm_batch_time / len(llm_batch):.3f}s/memory"
|
||||
)
|
||||
|
||||
# Build summary
|
||||
@@ -298,6 +316,10 @@ async def run_consolidation_job(
|
||||
if "db_write" in perf.timings:
|
||||
timing_parts.append(f"db_write={perf.timings['db_write']:.3f}s")
|
||||
|
||||
if perf.llm_calls > 0:
|
||||
timing_parts.append(f"avg_obs={perf.total_obs_in_context / perf.llm_calls:.1f}")
|
||||
timing_parts.append(f"avg_prompt_tokens=~{perf.total_prompt_chars / perf.llm_calls / 4:.0f}")
|
||||
|
||||
if timing_parts:
|
||||
perf.log(f"[4] Timing breakdown: {', '.join(timing_parts)}")
|
||||
|
||||
@@ -411,211 +433,219 @@ async def _trigger_mental_model_refreshes(
|
||||
return refreshed_count
|
||||
|
||||
|
||||
async def _process_memory(
|
||||
async def _process_memory_batch(
|
||||
conn: "Connection",
|
||||
memory_engine: "MemoryEngine",
|
||||
bank_id: str,
|
||||
memory: dict[str, Any],
|
||||
memories: list[dict[str, Any]],
|
||||
request_context: "RequestContext",
|
||||
perf: ConsolidationPerfLog | None = None,
|
||||
config: Any = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Process a single memory for consolidation using a SINGLE LLM call.
|
||||
Process a batch of memories in a single LLM call.
|
||||
|
||||
This function:
|
||||
1. Finds related observations (can be empty)
|
||||
2. Uses ONE LLM call to extract durable knowledge AND decide on actions
|
||||
3. Executes array of actions (can be multiple creates/updates)
|
||||
Steps:
|
||||
1. Parallel recalls — one per fact (read-only; safe to parallelise)
|
||||
2. Union of retrieved observations across the batch (deduped by id)
|
||||
3. Single LLM call with all N facts + unioned observations
|
||||
4. Sequential action execution (writes remain serial for consistency)
|
||||
5. Returns one result dict per memory, in the same order as `memories`
|
||||
|
||||
The LLM handles all cases:
|
||||
- No related observations: returns create action(s) with extracted durable knowledge
|
||||
- Related observations exist: returns update/create actions based on tag routing
|
||||
- Purely ephemeral fact: returns empty array (skip)
|
||||
|
||||
Returns:
|
||||
Dict with action summary: created/updated/merged counts
|
||||
Per-fact security: action execution validates each learning_id against the
|
||||
observations that were recalled specifically for that fact, so cross-tag
|
||||
updates cannot occur.
|
||||
"""
|
||||
from ...tracing import get_tracer, is_tracing_enabled
|
||||
import asyncio
|
||||
|
||||
fact_text = memory["text"]
|
||||
memory_id = memory["id"]
|
||||
fact_tags = memory.get("tags") or []
|
||||
|
||||
# Create parent span for this memory's consolidation
|
||||
tracer = get_tracer()
|
||||
if is_tracing_enabled():
|
||||
consolidation_span = tracer.start_span("hindsight.consolidation")
|
||||
consolidation_span.set_attribute("hindsight.memory_id", str(memory_id))
|
||||
consolidation_span.set_attribute("hindsight.bank_id", bank_id)
|
||||
else:
|
||||
consolidation_span = None
|
||||
|
||||
try:
|
||||
# Find related observations using the full recall system
|
||||
# SECURITY: Pass tags to ensure observations don't leak across security boundaries
|
||||
t0 = time.time()
|
||||
recall_result = await _find_related_observations(
|
||||
# 1. Parallel recalls — one per fact
|
||||
t0 = time.time()
|
||||
recall_tasks = [
|
||||
_find_related_observations(
|
||||
memory_engine=memory_engine,
|
||||
bank_id=bank_id,
|
||||
query=fact_text,
|
||||
query=m["text"],
|
||||
request_context=request_context,
|
||||
tags=fact_tags, # Pass source memory's tags for security
|
||||
tags=m.get("tags") or [],
|
||||
)
|
||||
if perf:
|
||||
perf.record_timing("recall", time.time() - t0)
|
||||
for m in memories
|
||||
]
|
||||
per_fact_recalls = await asyncio.gather(*recall_tasks)
|
||||
if perf:
|
||||
perf.record_timing("recall", time.time() - t0)
|
||||
|
||||
# Single LLM call handles ALL cases (with or without existing observations)
|
||||
# Note: Tags are NOT passed to LLM - they are handled algorithmically
|
||||
t0 = time.time()
|
||||
actions = await _consolidate_with_llm(
|
||||
# 2. Build per-fact observation sets (keyed by memory ID string) for secure action validation
|
||||
per_fact_obs_ids: dict[str, set[str]] = {
|
||||
str(memories[i]["id"]): {str(obs.id) for obs in r.results} for i, r in enumerate(per_fact_recalls)
|
||||
}
|
||||
|
||||
# Union all observations (deduped by id)
|
||||
seen_ids: set[str] = set()
|
||||
union_observations: list["MemoryFact"] = []
|
||||
union_source_facts: dict[str, "MemoryFact"] = {}
|
||||
for recall_result in per_fact_recalls:
|
||||
for obs in recall_result.results:
|
||||
obs_id = str(obs.id)
|
||||
if obs_id not in seen_ids:
|
||||
seen_ids.add(obs_id)
|
||||
union_observations.append(obs)
|
||||
if recall_result.source_facts:
|
||||
union_source_facts.update(recall_result.source_facts)
|
||||
|
||||
# 3. Single LLM call
|
||||
t0 = time.time()
|
||||
llm_result = await _consolidate_batch_with_llm(
|
||||
memory_engine=memory_engine,
|
||||
memories=memories,
|
||||
union_observations=union_observations,
|
||||
union_source_facts=union_source_facts,
|
||||
config=config,
|
||||
)
|
||||
if perf:
|
||||
perf.record_timing("llm", time.time() - t0)
|
||||
perf.record_llm_call(llm_result.obs_count, llm_result.prompt_chars)
|
||||
|
||||
# 4. Sequential execution of creates / updates / deletes
|
||||
# Track which memory indices participated so we can build per-memory results for stats
|
||||
per_memory_created: set[str] = set()
|
||||
per_memory_updated: set[str] = set()
|
||||
|
||||
# All memories in the batch share the same tag set (enforced by batching)
|
||||
fact_tags = memories[0].get("tags") or [] if memories else []
|
||||
|
||||
mem_by_id = {str(m["id"]): m for m in memories}
|
||||
|
||||
for create in llm_result.creates:
|
||||
source_mems = [mem_by_id[fid] for fid in create.source_fact_ids if fid in mem_by_id]
|
||||
if not source_mems:
|
||||
continue
|
||||
await _execute_create_action(
|
||||
conn=conn,
|
||||
memory_engine=memory_engine,
|
||||
fact_text=fact_text,
|
||||
recall_result=recall_result,
|
||||
config=config,
|
||||
bank_id=bank_id,
|
||||
source_memory_ids=[m["id"] for m in source_mems],
|
||||
text=create.text,
|
||||
source_fact_tags=fact_tags,
|
||||
event_date=_min_date(m.get("event_date") for m in source_mems),
|
||||
occurred_start=_min_date(m.get("occurred_start") for m in source_mems),
|
||||
occurred_end=_max_date(m.get("occurred_end") for m in source_mems),
|
||||
mentioned_at=_max_date(m.get("mentioned_at") for m in source_mems),
|
||||
perf=perf,
|
||||
)
|
||||
if perf:
|
||||
perf.record_timing("llm", time.time() - t0)
|
||||
for m in source_mems:
|
||||
per_memory_created.add(str(m["id"]))
|
||||
|
||||
if not actions:
|
||||
# LLM returned empty array - fact is purely ephemeral, skip
|
||||
return {"action": "skipped", "reason": "no_durable_knowledge"}
|
||||
for update in llm_result.updates:
|
||||
source_mems = [mem_by_id[fid] for fid in update.source_fact_ids if fid in mem_by_id]
|
||||
if not source_mems:
|
||||
continue
|
||||
# Security: the observation must have been recalled for at least one of the source facts
|
||||
if not any(update.observation_id in per_fact_obs_ids.get(str(m["id"]), set()) for m in source_mems):
|
||||
logger.debug(
|
||||
f"Batch consolidation: rejected update — observation {update.observation_id} "
|
||||
f"not in any source fact's recall"
|
||||
)
|
||||
continue
|
||||
await _execute_update_action(
|
||||
conn=conn,
|
||||
memory_engine=memory_engine,
|
||||
bank_id=bank_id,
|
||||
source_memory_ids=[m["id"] for m in source_mems],
|
||||
observation_id=update.observation_id,
|
||||
new_text=update.text,
|
||||
observations=union_observations,
|
||||
source_fact_tags=fact_tags,
|
||||
source_occurred_start=_min_date(m.get("occurred_start") for m in source_mems),
|
||||
source_occurred_end=_max_date(m.get("occurred_end") for m in source_mems),
|
||||
source_mentioned_at=_max_date(m.get("mentioned_at") for m in source_mems),
|
||||
perf=perf,
|
||||
)
|
||||
for m in source_mems:
|
||||
per_memory_updated.add(str(m["id"]))
|
||||
|
||||
# Execute all actions and collect results
|
||||
results = []
|
||||
for action in actions:
|
||||
action_type = action.get("action")
|
||||
if action_type == "update":
|
||||
result = await _execute_update_action(
|
||||
conn=conn,
|
||||
memory_engine=memory_engine,
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
action=action,
|
||||
observations=recall_result.results,
|
||||
source_fact_tags=fact_tags, # Pass source fact's tags for security
|
||||
source_occurred_start=memory.get("occurred_start"),
|
||||
source_occurred_end=memory.get("occurred_end"),
|
||||
source_mentioned_at=memory.get("mentioned_at"),
|
||||
perf=perf,
|
||||
)
|
||||
results.append(result)
|
||||
elif action_type == "create":
|
||||
result = await _execute_create_action(
|
||||
conn=conn,
|
||||
memory_engine=memory_engine,
|
||||
bank_id=bank_id,
|
||||
memory_id=memory_id,
|
||||
action=action,
|
||||
source_fact_tags=fact_tags, # Pass source fact's tags for security
|
||||
event_date=memory.get("event_date"),
|
||||
occurred_start=memory.get("occurred_start"),
|
||||
occurred_end=memory.get("occurred_end"),
|
||||
mentioned_at=memory.get("mentioned_at"),
|
||||
perf=perf,
|
||||
)
|
||||
results.append(result)
|
||||
for delete in llm_result.deletes:
|
||||
# Security: the observation must be present in the unioned recall
|
||||
if not any(str(obs.id) == delete.observation_id for obs in union_observations):
|
||||
logger.debug(
|
||||
f"Batch consolidation: rejected delete — observation {delete.observation_id} not in unioned recall"
|
||||
)
|
||||
continue
|
||||
await _execute_delete_action(conn=conn, bank_id=bank_id, observation_id=delete.observation_id)
|
||||
|
||||
if not results:
|
||||
# No valid actions executed
|
||||
return {"action": "skipped", "reason": "no_valid_actions"}
|
||||
# Build per-memory result dicts for the stats tracker in the outer loop
|
||||
results: list[dict[str, Any]] = []
|
||||
for m in memories:
|
||||
mid = str(m["id"])
|
||||
created = mid in per_memory_created
|
||||
updated = mid in per_memory_updated
|
||||
if created and updated:
|
||||
results.append({"action": "multiple", "created": 1, "updated": 1, "merged": 0, "total_actions": 2})
|
||||
elif created:
|
||||
results.append({"action": "created"})
|
||||
elif updated:
|
||||
results.append({"action": "updated"})
|
||||
else:
|
||||
results.append({"action": "skipped", "reason": "no_durable_knowledge"})
|
||||
|
||||
# Summarize results
|
||||
created = sum(1 for r in results if r.get("action") == "created")
|
||||
updated = sum(1 for r in results if r.get("action") == "updated")
|
||||
merged = sum(1 for r in results if r.get("action") == "merged")
|
||||
return results
|
||||
|
||||
if len(results) == 1:
|
||||
return results[0]
|
||||
|
||||
return {
|
||||
"action": "multiple",
|
||||
"created": created,
|
||||
"updated": updated,
|
||||
"merged": merged,
|
||||
"total_actions": len(results),
|
||||
}
|
||||
finally:
|
||||
if consolidation_span:
|
||||
consolidation_span.end()
|
||||
def _min_date(dates: "Any") -> "datetime | None":
|
||||
"""Return the minimum non-None datetime from an iterable."""
|
||||
return min((d for d in dates if d is not None), default=None)
|
||||
|
||||
|
||||
def _max_date(dates: "Any") -> "datetime | None":
|
||||
"""Return the maximum non-None datetime from an iterable."""
|
||||
return max((d for d in dates if d is not None), default=None)
|
||||
|
||||
|
||||
async def _execute_update_action(
|
||||
conn: "Connection",
|
||||
memory_engine: "MemoryEngine",
|
||||
bank_id: str,
|
||||
memory_id: uuid.UUID,
|
||||
action: dict[str, Any],
|
||||
source_memory_ids: list[uuid.UUID],
|
||||
observation_id: str,
|
||||
new_text: str,
|
||||
observations: list["MemoryFact"],
|
||||
source_fact_tags: list[str] | None = None,
|
||||
source_occurred_start: datetime | None = None,
|
||||
source_occurred_end: datetime | None = None,
|
||||
source_mentioned_at: datetime | None = None,
|
||||
perf: ConsolidationPerfLog | None = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> None:
|
||||
"""
|
||||
Execute an update action on an existing observation.
|
||||
Update an existing observation.
|
||||
|
||||
Updates the observation text, adds to history, increments proof_count,
|
||||
and updates temporal fields:
|
||||
- occurred_start: uses LEAST to keep the earliest start time
|
||||
- occurred_end: uses GREATEST to keep the most recent end time
|
||||
- mentioned_at: uses GREATEST to keep the most recent mention time
|
||||
|
||||
SECURITY: Merges source fact's tags into the observation's existing tags.
|
||||
This ensures all contributors can see the observation they contributed to.
|
||||
For example, if Lisa's observation (tags=['user_lisa']) is updated with
|
||||
Mike's fact (tags=['user_mike']), the observation will have both tags.
|
||||
Extends source_memory_ids with all contributing memories, updates temporal fields
|
||||
(LEAST for occurred_start, GREATEST for occurred_end / mentioned_at), and merges tags.
|
||||
"""
|
||||
learning_id = action.get("learning_id")
|
||||
new_text = action.get("text")
|
||||
reason = action.get("reason", "Updated with new fact")
|
||||
|
||||
if not learning_id or not new_text:
|
||||
return {"action": "skipped", "reason": "missing_learning_id_or_text"}
|
||||
|
||||
# Find the observation
|
||||
model = next((m for m in observations if m.id == learning_id), None)
|
||||
model = next((m for m in observations if str(m.id) == observation_id), None)
|
||||
if not model:
|
||||
return {"action": "skipped", "reason": "learning_not_found"}
|
||||
logger.debug(f"Update skipped: observation {observation_id} not found in recall results")
|
||||
return
|
||||
|
||||
# Build history entry (history is fetched fresh from DB on update to avoid stale state)
|
||||
history = [
|
||||
{
|
||||
"previous_text": model.text,
|
||||
"changed_at": datetime.now(timezone.utc).isoformat(),
|
||||
"reason": reason,
|
||||
"source_memory_id": str(memory_id),
|
||||
"source_memory_ids": [str(mid) for mid in source_memory_ids],
|
||||
}
|
||||
]
|
||||
|
||||
# Update source_memory_ids
|
||||
source_ids = list(model.source_fact_ids or [])
|
||||
source_ids.append(memory_id)
|
||||
source_ids = list(model.source_fact_ids or []) + source_memory_ids
|
||||
|
||||
# SECURITY: Merge source fact's tags into existing observation tags
|
||||
# This ensures all contributors can see the observation they contributed to
|
||||
# SECURITY: Merge source fact's tags into existing observation tags so all contributors can see it
|
||||
existing_tags = set(model.tags or [])
|
||||
source_tags = set(source_fact_tags or [])
|
||||
merged_tags = list(existing_tags | source_tags) # Union of both tag sets
|
||||
if source_tags and source_tags != existing_tags:
|
||||
logger.debug(
|
||||
f"Security: Merging tags for observation {learning_id}: "
|
||||
f"existing={list(existing_tags)}, source={list(source_tags)}, merged={merged_tags}"
|
||||
)
|
||||
merged_tags = list(existing_tags | source_tags)
|
||||
|
||||
# Generate new embedding for updated text
|
||||
t0 = time.time()
|
||||
embeddings = await embedding_utils.generate_embeddings_batch(memory_engine.embeddings, [new_text])
|
||||
embedding_str = str(embeddings[0]) if embeddings else None
|
||||
if perf:
|
||||
perf.record_timing("embedding", time.time() - t0)
|
||||
|
||||
# Update the observation
|
||||
# - occurred_start: LEAST keeps the earliest start time across all source facts
|
||||
# - occurred_end: GREATEST keeps the most recent end time across all source facts
|
||||
# - mentioned_at: GREATEST keeps the most recent mention time
|
||||
# - tags: merged from existing + source fact (for visibility)
|
||||
t0 = time.time()
|
||||
await conn.execute(
|
||||
f"""
|
||||
@@ -637,73 +667,65 @@ async def _execute_update_action(
|
||||
json.dumps(history),
|
||||
source_ids,
|
||||
len(source_ids),
|
||||
uuid.UUID(learning_id),
|
||||
uuid.UUID(observation_id),
|
||||
source_occurred_start,
|
||||
source_occurred_end,
|
||||
source_mentioned_at,
|
||||
merged_tags,
|
||||
)
|
||||
|
||||
# Create links from memory to observation
|
||||
await _create_memory_links(conn, memory_id, uuid.UUID(learning_id))
|
||||
if perf:
|
||||
perf.record_timing("db_write", time.time() - t0)
|
||||
|
||||
logger.debug(f"Updated observation {learning_id} with memory {memory_id}")
|
||||
|
||||
return {"action": "updated", "observation_id": learning_id}
|
||||
logger.debug(f"Updated observation {observation_id} from {len(source_memory_ids)} source memories")
|
||||
|
||||
|
||||
async def _execute_create_action(
|
||||
conn: "Connection",
|
||||
memory_engine: "MemoryEngine",
|
||||
bank_id: str,
|
||||
memory_id: uuid.UUID,
|
||||
action: dict[str, Any],
|
||||
source_memory_ids: list[uuid.UUID],
|
||||
text: str,
|
||||
source_fact_tags: list[str] | None = None,
|
||||
event_date: datetime | None = None,
|
||||
occurred_start: datetime | None = None,
|
||||
occurred_end: datetime | None = None,
|
||||
mentioned_at: datetime | None = None,
|
||||
perf: ConsolidationPerfLog | None = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> None:
|
||||
"""
|
||||
Execute a create action for a new observation.
|
||||
Create a new observation from one or more source memories.
|
||||
|
||||
Creates a new observation with the specified text.
|
||||
The text comes directly from the classify LLM - no second LLM call needed.
|
||||
|
||||
Tags are determined algorithmically (not by LLM):
|
||||
- Observations always inherit their source fact's tags
|
||||
- This ensures visibility scope is maintained (security)
|
||||
Tags are inherited from the source facts (determined algorithmically, not by LLM)
|
||||
to maintain visibility scope.
|
||||
"""
|
||||
text = action.get("text")
|
||||
|
||||
# Tags are determined algorithmically - always use source fact's tags
|
||||
# This ensures private memories create private observations
|
||||
tags = source_fact_tags or []
|
||||
|
||||
if not text:
|
||||
return {"action": "skipped", "reason": "missing_text"}
|
||||
|
||||
# Use text directly from classify - skip the redundant LLM call
|
||||
result = await _create_observation_directly(
|
||||
await _create_observation_directly(
|
||||
conn=conn,
|
||||
memory_engine=memory_engine,
|
||||
bank_id=bank_id,
|
||||
source_memory_id=memory_id,
|
||||
observation_text=text, # Text already processed by classify LLM
|
||||
tags=tags,
|
||||
source_memory_ids=source_memory_ids,
|
||||
observation_text=text,
|
||||
tags=source_fact_tags or [],
|
||||
event_date=event_date,
|
||||
occurred_start=occurred_start,
|
||||
occurred_end=occurred_end,
|
||||
mentioned_at=mentioned_at,
|
||||
perf=perf,
|
||||
)
|
||||
logger.debug(f"Created observation from {len(source_memory_ids)} source memories")
|
||||
|
||||
logger.debug(f"Created observation {result.get('observation_id')} from memory {memory_id} (tags: {tags})")
|
||||
|
||||
return result
|
||||
async def _execute_delete_action(
|
||||
conn: "Connection",
|
||||
bank_id: str,
|
||||
observation_id: str,
|
||||
) -> None:
|
||||
"""Delete a superseded or contradicted observation."""
|
||||
await conn.execute(
|
||||
f"DELETE FROM {fq_table('memory_units')} WHERE id = $1 AND bank_id = $2 AND fact_type = 'observation'",
|
||||
uuid.UUID(observation_id),
|
||||
bank_id,
|
||||
)
|
||||
logger.debug(f"Deleted observation {observation_id}")
|
||||
|
||||
|
||||
async def _create_memory_links(
|
||||
@@ -803,7 +825,6 @@ def _build_observations_for_llm(
|
||||
"id": obs.id,
|
||||
"text": obs.text,
|
||||
"proof_count": len(obs.source_fact_ids or []) or 1,
|
||||
"tags": obs.tags or [],
|
||||
}
|
||||
if obs.occurred_start:
|
||||
obs_data["occurred_start"] = obs.occurred_start
|
||||
@@ -811,74 +832,91 @@ def _build_observations_for_llm(
|
||||
obs_data["occurred_end"] = obs.occurred_end
|
||||
if obs.mentioned_at:
|
||||
obs_data["mentioned_at"] = obs.mentioned_at
|
||||
source_memories = [
|
||||
{"text": sf.text, "occurred_start": sf.occurred_start}
|
||||
for sid in (obs.source_fact_ids or [])[:3]
|
||||
if (sf := source_facts.get(sid)) is not None
|
||||
]
|
||||
source_memories = []
|
||||
for sid in obs.source_fact_ids or []:
|
||||
sf = source_facts.get(sid)
|
||||
if sf is None:
|
||||
continue
|
||||
sf_data: dict[str, Any] = {"text": sf.text}
|
||||
if sf.context:
|
||||
sf_data["context"] = sf.context
|
||||
if sf.occurred_start:
|
||||
sf_data["occurred_start"] = sf.occurred_start
|
||||
if sf.occurred_end:
|
||||
sf_data["occurred_end"] = sf.occurred_end
|
||||
if sf.mentioned_at:
|
||||
sf_data["mentioned_at"] = sf.mentioned_at
|
||||
source_memories.append(sf_data)
|
||||
if source_memories:
|
||||
obs_data["source_memories"] = source_memories
|
||||
obs_list.append(obs_data)
|
||||
return obs_list
|
||||
|
||||
|
||||
async def _consolidate_with_llm(
|
||||
async def _consolidate_batch_with_llm(
|
||||
memory_engine: "MemoryEngine",
|
||||
fact_text: str,
|
||||
recall_result: "RecallResult",
|
||||
memories: list[dict[str, Any]],
|
||||
union_observations: "list[MemoryFact]",
|
||||
union_source_facts: "dict[str, MemoryFact]",
|
||||
config: Any = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Single LLM call to extract durable knowledge and decide on consolidation actions.
|
||||
|
||||
This handles ALL cases:
|
||||
- No related observations: extracts durable knowledge, returns create action
|
||||
- Related observations exist: compares and returns update/create actions
|
||||
- Purely ephemeral fact: returns empty array
|
||||
|
||||
Note: Tags are NOT handled by the LLM. They are determined algorithmically:
|
||||
- CREATE: observation inherits source fact's tags
|
||||
- UPDATE: observation merges source fact's tags with existing tags
|
||||
|
||||
Returns:
|
||||
List of actions, each being:
|
||||
- {"action": "update", "learning_id": "uuid", "text": "...", "reason": "..."}
|
||||
- {"action": "create", "text": "...", "reason": "..."}
|
||||
- [] if fact is purely ephemeral (no durable knowledge)
|
||||
"""
|
||||
observations = recall_result.results
|
||||
source_facts = recall_result.source_facts or {}
|
||||
|
||||
if observations:
|
||||
obs_list = _build_observations_for_llm(observations, source_facts)
|
||||
) -> _BatchLLMResult:
|
||||
"""Single LLM call for a batch of facts against a pooled set of observations."""
|
||||
if union_observations:
|
||||
obs_list = _build_observations_for_llm(union_observations, union_source_facts)
|
||||
observations_text = json.dumps(obs_list, indent=2)
|
||||
else:
|
||||
observations_text = "[]"
|
||||
|
||||
def _fact_line(m: dict[str, Any]) -> str:
|
||||
parts = [f"[{m['id']}] {m['text']}"]
|
||||
if m.get("occurred_start"):
|
||||
parts.append(f"occurred_start={m['occurred_start']}")
|
||||
if m.get("occurred_end"):
|
||||
parts.append(f"occurred_end={m['occurred_end']}")
|
||||
if m.get("mentioned_at"):
|
||||
parts.append(f"mentioned_at={m['mentioned_at']}")
|
||||
return " | ".join(parts)
|
||||
|
||||
facts_lines = "\n".join(_fact_line(m) for m in memories)
|
||||
|
||||
observations_mission = config.observations_mission if config is not None else None
|
||||
prompt_template = build_consolidation_prompt(observations_mission)
|
||||
prompt_template = build_batch_consolidation_prompt(observations_mission)
|
||||
prompt = prompt_template.format(
|
||||
fact_text=fact_text,
|
||||
facts_text=facts_lines,
|
||||
observations_text=observations_text,
|
||||
)
|
||||
|
||||
messages = [
|
||||
{"role": "user", "content": prompt},
|
||||
]
|
||||
max_attempts = 3
|
||||
last_exc: Exception | None = None
|
||||
for attempt in range(1, max_attempts + 1):
|
||||
try:
|
||||
response: _ConsolidationBatchResponse = await memory_engine._consolidation_llm_config.call(
|
||||
messages=[{"role": "user", "content": prompt}],
|
||||
response_format=_ConsolidationBatchResponse,
|
||||
scope="consolidation",
|
||||
)
|
||||
return _BatchLLMResult(
|
||||
creates=response.creates,
|
||||
updates=response.updates,
|
||||
deletes=response.deletes,
|
||||
obs_count=len(union_observations),
|
||||
prompt_chars=len(prompt),
|
||||
)
|
||||
except Exception as exc:
|
||||
last_exc = exc
|
||||
logger.warning(f"[CONSOLIDATION] LLM batch call failed (attempt {attempt}/{max_attempts}): {exc}")
|
||||
|
||||
response: _ConsolidationResponse = await memory_engine._consolidation_llm_config.call(
|
||||
messages=messages,
|
||||
response_format=_ConsolidationResponse,
|
||||
scope="consolidation",
|
||||
logger.error(
|
||||
f"[CONSOLIDATION] LLM batch call failed after {max_attempts} attempts, skipping batch. Last error: {last_exc}"
|
||||
)
|
||||
return [a.model_dump() for a in response.actions]
|
||||
return _BatchLLMResult(obs_count=len(union_observations), prompt_chars=len(prompt))
|
||||
|
||||
|
||||
async def _create_observation_directly(
|
||||
conn: "Connection",
|
||||
memory_engine: "MemoryEngine",
|
||||
bank_id: str,
|
||||
source_memory_id: uuid.UUID,
|
||||
source_memory_ids: list[uuid.UUID],
|
||||
observation_text: str,
|
||||
tags: list[str] | None = None,
|
||||
event_date: datetime | None = None,
|
||||
@@ -887,12 +925,7 @@ async def _create_observation_directly(
|
||||
mentioned_at: datetime | None = None,
|
||||
perf: ConsolidationPerfLog | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Create an observation directly with pre-processed text (no LLM call).
|
||||
|
||||
Used when the classify LLM has already provided the learning text.
|
||||
This avoids the redundant second LLM call.
|
||||
"""
|
||||
"""Create an observation from one or more source memories with pre-processed text."""
|
||||
# Generate embedding for the observation (convert to string for pgvector)
|
||||
t0 = time.time()
|
||||
embeddings = await embedding_utils.generate_embeddings_batch(memory_engine.embeddings, [observation_text])
|
||||
@@ -942,7 +975,7 @@ async def _create_observation_directly(
|
||||
bank_id,
|
||||
observation_text,
|
||||
embedding_str,
|
||||
[source_memory_id],
|
||||
source_memory_ids,
|
||||
obs_tags,
|
||||
obs_event_date,
|
||||
obs_occurred_start,
|
||||
@@ -950,11 +983,9 @@ async def _create_observation_directly(
|
||||
obs_mentioned_at,
|
||||
)
|
||||
|
||||
# Create links between memory and observation (includes entity links, memory_links)
|
||||
await _create_memory_links(conn, source_memory_id, observation_id)
|
||||
if perf:
|
||||
perf.record_timing("db_write", time.time() - t0)
|
||||
|
||||
logger.debug(f"Created observation {observation_id} from memory {source_memory_id} (tags: {obs_tags})")
|
||||
logger.debug(f"Created observation {observation_id} from {len(source_memory_ids)} memories (tags: {obs_tags})")
|
||||
|
||||
return {"action": "created", "observation_id": str(row["id"]), "tags": obs_tags}
|
||||
|
||||
@@ -1,21 +1,21 @@
|
||||
"""Prompts for the consolidation engine."""
|
||||
|
||||
# Output format instructions
|
||||
_OUTPUT_FORMAT = """
|
||||
Output a JSON object with an "actions" array:
|
||||
{{"actions": [
|
||||
{{"action": "update", "learning_id": "uuid-from-observations", "text": "...", "reason": "..."}},
|
||||
{{"action": "create", "text": "...", "reason": "..."}}
|
||||
]}}
|
||||
# Default mission when no bank-specific mission is set
|
||||
_DEFAULT_MISSION = "Track every detail: names, numbers, dates, places, and relationships. Prefer specifics over abstractions, never generalise."
|
||||
|
||||
Return {{"actions": []}} if the fact contains no durable knowledge.
|
||||
Do NOT include "tags" in output — tags are handled automatically."""
|
||||
# Processing rules — always present regardless of mission
|
||||
_PROCESSING_RULES = """Processing rules (always apply):
|
||||
- REDUNDANT: same info worded differently → UPDATE the existing observation.
|
||||
- CONTRADICTION/UPDATE: capture both states with temporal markers ("used to X, now Y").
|
||||
- RESOLVE REFERENCES: when a new fact provides a concrete value resolving a vague placeholder in an existing observation (e.g. "home country", "hometown", "birthplace", "native language", "her ex", "that city"), UPDATE the observation to embed the resolved value explicitly. Example: new fact says "grandma in Sweden" + existing observation says "moved from her home country" → update to "home country is Sweden".
|
||||
- NEVER merge observations about different people or unrelated topics."""
|
||||
|
||||
# Data section - holds the dynamic per-call data
|
||||
_DATA_SECTION = """
|
||||
NEW FACT: {fact_text}
|
||||
# Data section — format placeholders {facts_text} and {observations_text} are substituted at call time
|
||||
_BATCH_DATA_SECTION = """
|
||||
NEW FACTS:
|
||||
{facts_text}
|
||||
|
||||
EXISTING OBSERVATIONS (JSON array with source memories and dates):
|
||||
EXISTING OBSERVATIONS (JSON array, pooled from recalls across all facts above):
|
||||
{observations_text}
|
||||
|
||||
Each observation includes:
|
||||
@@ -25,34 +25,42 @@ Each observation includes:
|
||||
- occurred_start/occurred_end: temporal range of source facts
|
||||
- source_memories: array of supporting facts with their text and dates
|
||||
|
||||
Compare the new fact against existing observations:
|
||||
- Same topic → UPDATE with learning_id
|
||||
- New topic → CREATE new observation
|
||||
- Purely ephemeral → return empty actions list"""
|
||||
Compare the facts against existing observations:
|
||||
- Same topic as an existing observation → UPDATE it (observation_id + source_fact_ids)
|
||||
- New topic with durable knowledge → CREATE a new observation (source_fact_ids)
|
||||
- Cross-reference facts within the batch: a later fact may resolve a vague reference in an earlier one
|
||||
- Purely ephemeral facts → omit them (no create/update needed)"""
|
||||
|
||||
# Default rules used when no observations_mission is set
|
||||
_DEFAULT_RULES = """Extract DURABLE KNOWLEDGE from facts — the stable truth implied by an event, not transient state.
|
||||
# Output format — JSON braces escaped as {{ }} so .format() leaves them literal
|
||||
_BATCH_OUTPUT_FORMAT = """
|
||||
Output a JSON object with three arrays.
|
||||
|
||||
Example: "User moved to Room 203" → observe "Room 203 exists", not "User is in Room 203".
|
||||
Example (showing the required UUID format for all IDs):
|
||||
{{"creates": [{{"text": "Alice lives in Berlin", "source_fact_ids": ["a1b2c3d4-e5f6-7890-abcd-ef1234567890", "b2c3d4e5-f6a7-8901-bcde-f12345678901"]}}],
|
||||
"updates": [{{"text": "Alice works at Acme Corp as a senior engineer", "observation_id": "c3d4e5f6-a7b8-9012-cdef-123456789012", "source_fact_ids": ["d4e5f6a7-b8c9-0123-defa-234567890123"]}}],
|
||||
"deletes": [{{"observation_id": "e5f6a7b8-c9d0-1234-efab-345678901234"}}]}}
|
||||
|
||||
Rules:
|
||||
- Keep specifics: names, numbers, locations. Never abstract into general principles.
|
||||
- NEVER merge observations about different people or unrelated topics.
|
||||
- REDUNDANT: same info worded differently → update existing.
|
||||
- CONTRADICTION/UPDATE: capture both states with temporal markers ("used to X, now Y").
|
||||
- RESOLVE REFERENCES: When a new fact provides a concrete value that resolves a vague placeholder in an existing observation (e.g., a location that corresponds to "home country", "hometown", "birthplace", "native language", "her ex", "that city"), UPDATE the existing observation to embed the resolved value explicitly. Example: new fact mentions grandma in Sweden + existing observation says "moved from her home country" → update to state "home country is Sweden"."""
|
||||
- "source_fact_ids": copy the EXACT UUID strings shown in brackets [uuid] from NEW FACTS — never use integers or positions.
|
||||
- "observation_id": copy the EXACT "id" UUID string from EXISTING OBSERVATIONS.
|
||||
- One create/update may reference multiple facts when they jointly support the observation.
|
||||
- "deletes": only when an observation is directly superseded or contradicted by new facts.
|
||||
- Do NOT include "tags" — handled automatically.
|
||||
- Return {{"creates": [], "updates": [], "deletes": []}} if nothing durable is found."""
|
||||
|
||||
|
||||
def build_consolidation_prompt(observations_mission: str | None = None) -> str:
|
||||
def build_batch_consolidation_prompt(observations_mission: str | None = None) -> str:
|
||||
"""
|
||||
Build the consolidation prompt.
|
||||
Build the consolidation prompt for batch mode (multiple facts per LLM call).
|
||||
|
||||
If observations_mission is provided, it replaces the default durable-knowledge rules
|
||||
with bank-specific instructions for what to synthesise. Otherwise the default rules apply.
|
||||
The mission defines *what* to track (customisable per bank).
|
||||
Processing rules and output format are always present regardless of mission.
|
||||
"""
|
||||
rules_section = f"## MISSION\n{observations_mission}" if observations_mission else _DEFAULT_RULES
|
||||
mission = observations_mission or _DEFAULT_MISSION
|
||||
|
||||
return (
|
||||
"You are a memory consolidation system. Synthesize facts into observations "
|
||||
"and merge with existing observations when appropriate.\n\n" + rules_section + _DATA_SECTION + _OUTPUT_FORMAT
|
||||
"and merge with existing observations when appropriate.\n\n"
|
||||
f"## MISSION\n{mission}\n\n"
|
||||
f"{_PROCESSING_RULES}" + _BATCH_DATA_SECTION + _BATCH_OUTPUT_FORMAT
|
||||
)
|
||||
|
||||
@@ -273,6 +273,7 @@ def main():
|
||||
file_delete_after_retain=config.file_delete_after_retain,
|
||||
enable_observations=config.enable_observations,
|
||||
consolidation_batch_size=config.consolidation_batch_size,
|
||||
consolidation_llm_batch_size=config.consolidation_llm_batch_size,
|
||||
consolidation_max_tokens=config.consolidation_max_tokens,
|
||||
observations_mission=config.observations_mission,
|
||||
skip_llm_verification=config.skip_llm_verification,
|
||||
|
||||
@@ -1990,35 +1990,35 @@ class TestMentalModelRefreshAfterConsolidation:
|
||||
|
||||
|
||||
def test_consolidation_prompt_default():
|
||||
"""Test that the default consolidation prompt contains the built-in durable-knowledge rules."""
|
||||
from hindsight_api.engine.consolidation.prompts import build_consolidation_prompt
|
||||
"""Test that the default consolidation prompt contains the built-in mission and processing rules."""
|
||||
from hindsight_api.engine.consolidation.prompts import build_batch_consolidation_prompt
|
||||
|
||||
prompt = build_consolidation_prompt()
|
||||
assert "DURABLE KNOWLEDGE" in prompt
|
||||
prompt = build_batch_consolidation_prompt()
|
||||
assert "temporal markers" in prompt
|
||||
assert "{fact_text}" in prompt
|
||||
assert "RESOLVE REFERENCES" in prompt
|
||||
assert "{facts_text}" in prompt
|
||||
assert "{observations_text}" in prompt
|
||||
|
||||
|
||||
def test_consolidation_prompt_observations_mission():
|
||||
"""Test that observations_mission replaces the default rules."""
|
||||
from hindsight_api.engine.consolidation.prompts import build_consolidation_prompt
|
||||
"""Test that observations_mission replaces the default mission but keeps processing rules."""
|
||||
from hindsight_api.engine.consolidation.prompts import build_batch_consolidation_prompt
|
||||
|
||||
spec = "Observations are weekly summaries of sprint outcomes and team dynamics."
|
||||
prompt = build_consolidation_prompt(observations_mission=spec)
|
||||
prompt = build_batch_consolidation_prompt(observations_mission=spec)
|
||||
|
||||
# Spec is injected
|
||||
assert spec in prompt
|
||||
# Default rules are NOT present
|
||||
assert "EXTRACT DURABLE KNOWLEDGE" not in prompt
|
||||
# Output format and data placeholders remain
|
||||
assert "actions" in prompt
|
||||
assert "{fact_text}" in prompt
|
||||
# Processing rules and output format always remain
|
||||
assert "RESOLVE REFERENCES" in prompt
|
||||
assert "creates" in prompt
|
||||
assert "updates" in prompt
|
||||
assert "{facts_text}" in prompt
|
||||
assert "{observations_text}" in prompt
|
||||
|
||||
# Renders cleanly
|
||||
rendered = prompt.format(fact_text="Alice fixed a bug.", observations_text="[]")
|
||||
assert "{fact_text}" not in rendered
|
||||
rendered = prompt.format(facts_text="Alice fixed a bug.", observations_text="[]")
|
||||
assert "{facts_text}" not in rendered
|
||||
assert spec in rendered
|
||||
|
||||
|
||||
|
||||
@@ -455,48 +455,23 @@ class BenchmarkRunner:
|
||||
Get the count of memories pending consolidation.
|
||||
|
||||
Returns:
|
||||
Number of memories not yet consolidated into mental models
|
||||
Number of memories not yet processed by the consolidation job
|
||||
"""
|
||||
pool = await self.memory._get_pool()
|
||||
from hindsight_api.engine.memory_engine import fq_table
|
||||
|
||||
async with pool.acquire() as conn:
|
||||
# Check when consolidation last ran
|
||||
last_consolidated_row = await conn.fetchrow(
|
||||
result = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT MAX(created_at) as last_consolidated_at
|
||||
SELECT COUNT(*) as count
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $1 AND fact_type = 'mental_model'
|
||||
WHERE bank_id = $1 AND consolidated_at IS NULL AND fact_type IN ('experience', 'world')
|
||||
""",
|
||||
bank_id,
|
||||
)
|
||||
last_consolidated_at = last_consolidated_row["last_consolidated_at"] if last_consolidated_row else None
|
||||
|
||||
if last_consolidated_at:
|
||||
# Count memories created after last consolidation
|
||||
result = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT COUNT(*) as count
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $1 AND fact_type IN ('experience', 'world')
|
||||
AND created_at > $2
|
||||
""",
|
||||
bank_id,
|
||||
last_consolidated_at,
|
||||
)
|
||||
else:
|
||||
# If never consolidated, count all experience/world memories
|
||||
result = await conn.fetchrow(
|
||||
f"""
|
||||
SELECT COUNT(*) as count
|
||||
FROM {fq_table("memory_units")}
|
||||
WHERE bank_id = $1 AND fact_type IN ('experience', 'world')
|
||||
""",
|
||||
bank_id,
|
||||
)
|
||||
return result["count"] if result else 0
|
||||
|
||||
async def _wait_for_consolidation(self, bank_id: str, poll_interval: float = 2.0, timeout: float = 300.0) -> None:
|
||||
async def _wait_for_consolidation(self, bank_id: str, poll_interval: float = 2.0, timeout: float = 3000.0) -> None:
|
||||
"""
|
||||
Wait for consolidation to complete (pending_consolidation reaches 0).
|
||||
|
||||
@@ -860,6 +835,7 @@ class BenchmarkRunner:
|
||||
question_semaphore: asyncio.Semaphore,
|
||||
eval_semaphore_size: int = 8,
|
||||
clear_this_agent: bool = True,
|
||||
wait_consolidation: bool = False,
|
||||
) -> Dict:
|
||||
"""
|
||||
Process a single item (ingest + evaluate).
|
||||
@@ -867,6 +843,7 @@ class BenchmarkRunner:
|
||||
Args:
|
||||
clear_this_agent: Whether to clear this agent's data before ingesting.
|
||||
Set to False to skip clearing (e.g., when agent_id is shared and already cleared)
|
||||
wait_consolidation: If True, wait for consolidation to complete before evaluating QA.
|
||||
|
||||
Returns:
|
||||
Result dict with metrics
|
||||
@@ -891,6 +868,12 @@ class BenchmarkRunner:
|
||||
else:
|
||||
num_sessions = -1
|
||||
|
||||
# Wait for consolidation before evaluating if requested
|
||||
if wait_consolidation:
|
||||
step += 1
|
||||
console.print(f" [{step}] Waiting for consolidation...")
|
||||
await self._wait_for_consolidation(agent_id)
|
||||
|
||||
# Evaluate QA
|
||||
step += 1
|
||||
qa_pairs = self.dataset.get_qa_pairs(item)
|
||||
@@ -934,6 +917,7 @@ class BenchmarkRunner:
|
||||
max_concurrent_items: int = 1, # Max concurrent items (conversations) to process in parallel
|
||||
output_path: Optional[Path] = None, # Path to save results incrementally
|
||||
merge_with_existing: bool = False, # Whether to merge with existing results
|
||||
wait_consolidation: bool = False, # Wait for consolidation to complete before evaluating QA
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Run the full benchmark evaluation.
|
||||
@@ -1011,6 +995,7 @@ class BenchmarkRunner:
|
||||
max_concurrent_items,
|
||||
output_path,
|
||||
merge_with_existing,
|
||||
wait_consolidation,
|
||||
)
|
||||
|
||||
async def _run_single_phase(
|
||||
@@ -1028,6 +1013,7 @@ class BenchmarkRunner:
|
||||
max_concurrent_items: int = 1,
|
||||
output_path: Optional[Path] = None,
|
||||
merge_with_existing: bool = False,
|
||||
wait_consolidation: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
"""Original single-phase approach: process each item independently."""
|
||||
# Create semaphore for question processing
|
||||
@@ -1049,6 +1035,7 @@ class BenchmarkRunner:
|
||||
max_concurrent_items,
|
||||
output_path,
|
||||
merge_with_existing,
|
||||
wait_consolidation,
|
||||
)
|
||||
else:
|
||||
# Sequential item processing (original behavior)
|
||||
@@ -1065,6 +1052,7 @@ class BenchmarkRunner:
|
||||
filln,
|
||||
output_path,
|
||||
merge_with_existing,
|
||||
wait_consolidation,
|
||||
)
|
||||
|
||||
# Calculate overall metrics
|
||||
@@ -1100,6 +1088,7 @@ class BenchmarkRunner:
|
||||
filln: bool,
|
||||
output_path: Optional[Path] = None,
|
||||
merge_with_existing: bool = False,
|
||||
wait_consolidation: bool = False,
|
||||
) -> List[Dict]:
|
||||
"""Process items sequentially (original behavior)."""
|
||||
all_results = []
|
||||
@@ -1147,6 +1136,7 @@ class BenchmarkRunner:
|
||||
question_semaphore,
|
||||
eval_semaphore_size,
|
||||
clear_this_agent,
|
||||
wait_consolidation,
|
||||
)
|
||||
|
||||
# Replace existing result or append new one
|
||||
@@ -1178,6 +1168,7 @@ class BenchmarkRunner:
|
||||
max_concurrent_items: int,
|
||||
output_path: Optional[Path] = None,
|
||||
merge_with_existing: bool = False,
|
||||
wait_consolidation: bool = False,
|
||||
) -> List[Dict]:
|
||||
"""Process items in parallel (requires unique agent IDs per item)."""
|
||||
# Load existing results if merge_with_existing is True
|
||||
@@ -1222,6 +1213,7 @@ class BenchmarkRunner:
|
||||
question_semaphore,
|
||||
eval_semaphore_size,
|
||||
clear_this_agent=True, # Always clear for parallel processing
|
||||
wait_consolidation=wait_consolidation,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@@ -295,6 +295,7 @@ async def run_benchmark(
|
||||
only_failed: bool = False,
|
||||
only_invalid: bool = False,
|
||||
question_index: int = None,
|
||||
wait_consolidation: bool = False,
|
||||
):
|
||||
"""
|
||||
Run the LoComo benchmark.
|
||||
@@ -461,6 +462,7 @@ async def run_benchmark(
|
||||
max_concurrent_items=concurrent_items,
|
||||
output_path=output_path, # Save results incrementally
|
||||
merge_with_existing=merge_with_existing,
|
||||
wait_consolidation=wait_consolidation,
|
||||
)
|
||||
|
||||
# Display results (final save already happened incrementally)
|
||||
@@ -591,6 +593,11 @@ if __name__ == "__main__":
|
||||
default=None,
|
||||
help="Run only the question at this 0-based index within each conversation (e.g., 11)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--wait-consolidation",
|
||||
action="store_true",
|
||||
help="Wait for consolidation to complete after ingestion (or immediately when using --skip-ingestion) before evaluating QA.",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -610,5 +617,6 @@ if __name__ == "__main__":
|
||||
only_failed=args.only_failed,
|
||||
only_invalid=args.only_invalid,
|
||||
question_index=args.question_index,
|
||||
wait_consolidation=args.wait_consolidation,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -733,6 +733,7 @@ Observations are consolidated knowledge synthesized from facts.
|
||||
| `HINDSIGHT_API_ENABLE_OBSERVATIONS` | Enable observation consolidation | `true` |
|
||||
| `HINDSIGHT_API_CONSOLIDATION_BATCH_SIZE` | Memories to load per batch (internal optimization) | `50` |
|
||||
| `HINDSIGHT_API_CONSOLIDATION_MAX_TOKENS` | Max tokens for recall when finding related observations during consolidation | `1024` |
|
||||
| `HINDSIGHT_API_CONSOLIDATION_LLM_BATCH_SIZE` | Number of facts sent to the LLM in a single consolidation call. Higher values reduce LLM calls and improve throughput at the cost of larger prompts. Set to `1` to disable batching. | `8` |
|
||||
| `HINDSIGHT_API_OBSERVATIONS_MISSION` | What this bank should synthesise into durable observations. Replaces the built-in consolidation rules — leave unset to use the server default. | - |
|
||||
|
||||
#### Customizing observations: when to use what
|
||||
|
||||
Reference in New Issue
Block a user