""" Temporal + Semantic + Entity Memory System for AI Agents. This implements a sophisticated memory architecture that combines: 1. Temporal links: Memories connected by time proximity 2. Semantic links: Memories connected by meaning/similarity 3. Entity links: Memories connected by shared entities (PERSON, ORG, etc.) 4. Spreading activation: Search through the graph with activation decay 5. Dynamic weighting: Recency and frequency-based importance """ import os from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Optional, Tuple import asyncpg from sentence_transformers import SentenceTransformer from dotenv import load_dotenv import asyncio import time from concurrent.futures import ProcessPoolExecutor import numpy as np import uuid import logging from .utils import ( extract_facts, calculate_recency_weight, calculate_frequency_weight, ) from .entity_resolver import EntityResolver def utcnow(): """Get current UTC time with timezone info.""" return datetime.now(timezone.utc) # Logger for memory system logger = logging.getLogger(__name__) # Global process pool for parallel embedding generation # Each process loads its own copy of the embedding model # This provides TRUE parallelism for CPU-bound embedding operations _PROCESS_POOL = None _EMBEDDING_MODEL_NAME = "BAAI/bge-small-en-v1.5" # Process-local model cache (one per worker process) _worker_model = None def _get_worker_model(): """Get or load the embedding model in worker process.""" global _worker_model if _worker_model is None: _worker_model = SentenceTransformer(_EMBEDDING_MODEL_NAME) return _worker_model def _encode_batch_worker(texts: List[str]) -> List[List[float]]: """ Worker function for process pool - encodes texts to embeddings. This function runs in a separate process and loads its own model. """ model = _get_worker_model() embeddings = model.encode(texts, convert_to_numpy=True, show_progress_bar=False) return [emb.tolist() for emb in embeddings] def _get_process_pool(): """Get or create the global process pool.""" global _PROCESS_POOL if _PROCESS_POOL is None: # Use 4 worker processes for true parallelism # Adjust based on your CPU cores (each process loads ~500MB model) _PROCESS_POOL = ProcessPoolExecutor(max_workers=4) return _PROCESS_POOL class TemporalSemanticMemory: """ Advanced memory system using temporal and semantic linking with PostgreSQL. """ def __init__( self, db_url: Optional[str] = None, embedding_model: str = "BAAI/bge-small-en-v1.5", ): """ Initialize the temporal + semantic memory system. Args: db_url: PostgreSQL connection URL (postgresql://user:pass@host:port/dbname) embedding_model: Name of the SentenceTransformer model to use """ load_dotenv() # Initialize PostgreSQL connection URL self.db_url = db_url or os.getenv("DATABASE_URL") if not self.db_url: raise ValueError( "Database URL not found. " "Set DATABASE_URL environment variable." ) # Connection pool (created lazily on first use) self._pool = None self._pool_lock = asyncio.Lock() # Initialize entity resolver (will be created with pool) self.entity_resolver = None # Initialize local embedding model (384 dimensions) logger.info(f"Loading embedding model: {embedding_model}...") self.embedding_model = SentenceTransformer(embedding_model) logger.info(f"Model loaded (embedding dim: {self.embedding_model.get_sentence_embedding_dimension()})") # Background queue for access count updates (to avoid blocking searches) self._access_count_queue = asyncio.Queue() self._access_count_worker_task = None self._shutdown_event = asyncio.Event() async def _access_count_worker(self): """Background worker that processes access count updates in batches.""" pool = self._pool # Pool is guaranteed to exist when worker starts while not self._shutdown_event.is_set(): try: # Collect updates for up to 1 second or 1000 items updates = {} deadline = asyncio.get_event_loop().time() + 1.0 while len(updates) < 1000 and asyncio.get_event_loop().time() < deadline: try: # Wait for items with short timeout remaining_time = max(0.1, deadline - asyncio.get_event_loop().time()) node_ids = await asyncio.wait_for( self._access_count_queue.get(), timeout=remaining_time ) # Deduplicate by adding to set for node_id in node_ids: updates[node_id] = True except asyncio.TimeoutError: break # Process batch if we have updates if updates: node_id_list = list(updates.keys()) try: # Convert string UUIDs to UUID type for faster matching uuid_list = [uuid.UUID(nid) for nid in node_id_list] async with pool.acquire() as conn: await conn.execute( "UPDATE memory_units SET access_count = access_count + 1 WHERE id = ANY($1::uuid[])", uuid_list ) except Exception as e: logger.error(f"Access count worker: Error updating access counts: {e}") except asyncio.CancelledError: break except Exception as e: logger.error(f"Access count worker: Unexpected error: {e}") await asyncio.sleep(1) # Backoff on error async def _get_pool(self) -> asyncpg.Pool: """Get or create the connection pool (lazy initialization).""" if self._pool is None: async with self._pool_lock: if self._pool is None: self._pool = await asyncpg.create_pool( self.db_url, min_size=2, max_size=10, command_timeout=60, statement_cache_size=0 # Disable prepared statement cache ) # Initialize entity resolver with pool if self.entity_resolver is None: self.entity_resolver = EntityResolver(self._pool) # Start access count worker (outside lock, after pool is created) if self._access_count_worker_task is None and self._pool is not None: self._access_count_worker_task = asyncio.create_task(self._access_count_worker()) return self._pool async def close(self): """Close the connection pool and shutdown background workers.""" # Signal shutdown to worker self._shutdown_event.set() # Cancel and wait for worker task if self._access_count_worker_task is not None: self._access_count_worker_task.cancel() try: await self._access_count_worker_task except asyncio.CancelledError: pass # Close pool if self._pool is not None: await self._pool.close() self._pool = None def _generate_embedding(self, text: str) -> List[float]: """ Generate embedding for text using local SentenceTransformer model. Args: text: Text to embed Returns: 384-dimensional embedding vector (bge-small-en-v1.5) """ try: embedding = self.embedding_model.encode(text, convert_to_numpy=True, show_progress_bar=False) return embedding.tolist() except Exception as e: raise Exception(f"Failed to generate embedding: {str(e)}") async def _generate_embeddings_batch(self, texts: List[str]) -> List[List[float]]: """ Generate embeddings for multiple texts using local model in parallel. Uses a ProcessPoolExecutor to achieve TRUE parallelism for CPU-bound embedding generation. Each worker process loads its own model copy. When multiple put_async calls run in parallel, each can generate embeddings concurrently in separate processes (no GIL contention). Args: texts: List of texts to embed Returns: List of 384-dimensional embeddings in same order as input texts """ try: # Run in process pool for true parallelism loop = asyncio.get_event_loop() pool = _get_process_pool() embeddings = await loop.run_in_executor( pool, _encode_batch_worker, texts ) return embeddings except Exception as e: raise Exception(f"Failed to generate batch embeddings: {str(e)}") async def _find_duplicate_facts_batch( self, conn, agent_id: str, texts: List[str], embeddings: List[List[float]], event_date: datetime, time_window_hours: int = 24, similarity_threshold: float = 0.95 ) -> List[bool]: """ Check which facts are duplicates using semantic similarity + temporal window. For each new fact, checks if a semantically similar fact already exists within the time window. Uses pgvector cosine similarity for efficiency. Args: conn: Database connection agent_id: Agent identifier texts: List of fact texts to check embeddings: Corresponding embeddings event_date: Event date for temporal filtering time_window_hours: Hours before/after event_date to search (default: 24) similarity_threshold: Minimum cosine similarity to consider duplicate (default: 0.95) Returns: List of booleans - True if fact is a duplicate (should skip), False if new """ is_duplicate = [] time_lower = event_date - timedelta(hours=time_window_hours) time_upper = event_date + timedelta(hours=time_window_hours) for text, embedding in zip(texts, embeddings): # Query for similar facts within time window # Convert embedding list to string for asyncpg vector type embedding_str = str(embedding) result = await conn.fetchrow( """ SELECT id, text, 1 - (embedding <=> $1::vector) AS similarity FROM memory_units WHERE agent_id = $2 AND event_date BETWEEN $3 AND $4 AND 1 - (embedding <=> $1::vector) > $5 ORDER BY similarity DESC LIMIT 1 """, embedding_str, agent_id, time_lower, time_upper, similarity_threshold ) if result: is_duplicate.append(True) else: is_duplicate.append(False) return is_duplicate def put( self, agent_id: str, content: str, context: str = "", event_date: Optional[datetime] = None, ) -> List[str]: """ Store content as memory units (synchronous wrapper). This is a synchronous wrapper around put_async() for convenience. For best performance, use put_async() directly. Args: agent_id: Unique identifier for the agent content: Text content to store context: Context about when/why this memory was formed event_date: When the event occurred (defaults to now) Returns: List of created unit IDs """ # Run async version synchronously return asyncio.run(self.put_async(agent_id, content, context, event_date)) async def put_async( self, agent_id: str, content: str, context: str = "", event_date: Optional[datetime] = None, document_id: Optional[str] = None, document_metadata: Optional[Dict[str, Any]] = None, upsert: bool = False, ) -> List[str]: """ Store content as memory units with temporal and semantic links (ASYNC version). This is a convenience wrapper around put_batch_async for a single content item. Args: agent_id: Unique identifier for the agent content: Text content to store context: Context about when/why this memory was formed event_date: When the event occurred (defaults to now) document_id: Optional document ID for tracking and upsert document_metadata: Optional metadata about the document upsert: If True and document_id exists, delete old units and create new ones Returns: List of created unit IDs """ # Use put_batch_async with a single item (avoids code duplication) result = await self.put_batch_async( agent_id=agent_id, contents=[{ "content": content, "context": context, "event_date": event_date }], document_id=document_id, document_metadata=document_metadata, upsert=upsert ) # Return the first (and only) list of unit IDs return result[0] if result else [] async def put_batch_async( self, agent_id: str, contents: List[Dict[str, Any]], document_id: Optional[str] = None, document_metadata: Optional[Dict[str, Any]] = None, upsert: bool = False, ) -> List[List[str]]: """ Store multiple content items as memory units in ONE batch operation. This is MUCH more efficient than calling put_async multiple times: - Extracts facts from all contents in parallel - Generates ALL embeddings in ONE batch - Does ALL database operations in ONE transaction Args: agent_id: Unique identifier for the agent contents: List of dicts with keys: - "content" (required): Text content to store - "context" (optional): Context about the memory - "event_date" (optional): When the event occurred document_id: Optional document ID for tracking and upsert document_metadata: Optional metadata about the document upsert: If True and document_id exists, delete old units and create new ones Returns: List of lists of unit IDs (one list per content item) Example: unit_ids = await memory.put_batch_async( agent_id="user123", contents=[ {"content": "Alice works at Google", "context": "conversation"}, {"content": "Bob loves Python", "context": "conversation"}, ], document_id="meeting-2024-01-15", upsert=True ) # Returns: [["unit-id-1"], ["unit-id-2"]] """ start_time = time.time() logger.debug(f"\n{'='*60}") logger.debug(f"PUT_BATCH_ASYNC START: {agent_id}") logger.debug(f"Batch size: {len(contents)} content items") logger.debug(f"{'='*60}") if not contents: return [] # Step 1: Extract facts from ALL contents in parallel step_start = time.time() # Create tasks for parallel fact extraction 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) 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]) # 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 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 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', [])) end_idx = current_idx + len(fact_dicts) content_boundaries.append((start_idx, end_idx)) current_idx = end_idx total_facts = len(all_fact_texts) if total_facts == 0: return [[] for _ in contents] # Step 2: Generate ALL embeddings in ONE batch (HUGE speedup!) step_start = time.time() all_embeddings = await self._generate_embeddings_batch(all_fact_texts) logger.debug(f"[2] Generate embeddings (parallel): {len(all_embeddings)} embeddings in {time.time() - step_start:.3f}s") # Step 3: Process everything in ONE database transaction pool = await self._get_pool() async with pool.acquire() as conn: async with conn.transaction(): try: # Handle document tracking and upsert if 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() # 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 {}) ) logger.debug(f"[3.2] Document '{document_id}' stored/updated") # Deduplication check for all facts step_start = time.time() all_is_duplicate = [] for sentence, embedding, fact_date in zip(all_fact_texts, all_embeddings, all_fact_dates): dup_flags = await self._find_duplicate_facts_batch( conn, agent_id, [sentence], [embedding], fact_date ) all_is_duplicate.extend(dup_flags) duplicates_filtered = sum(all_is_duplicate) new_facts = total_facts - duplicates_filtered logger.debug(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] 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] results = await conn.fetch( """ INSERT INTO memory_units (agent_id, document_id, text, context, embedding, event_date, access_count) SELECT * FROM unnest($1::text[], $2::text[], $3::text[], $4::text[], $5::vector[], $6::timestamptz[], $7::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, [0] * len(filtered_sentences) ) created_unit_ids = [str(row['id']) for row in results] logger.debug(f"[5] Batch insert units: {len(created_unit_ids)} units in {time.time() - step_start:.3f}s") # Process entities for ALL units 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 ) logger.debug(f"[6] Process entities (batched): {time.time() - step_start:.3f}s") # Create temporal links step_start = time.time() await self._create_temporal_links_batch_per_fact(conn, agent_id, created_unit_ids) logger.debug(f"[7] Batch create temporal links: {time.time() - step_start:.3f}s") # Create semantic links step_start = time.time() await self._create_semantic_links_batch(conn, agent_id, created_unit_ids, filtered_embeddings) logger.debug(f"[8] Batch create semantic links: {time.time() - step_start:.3f}s") # Insert entity links step_start = time.time() if all_entity_links: await self._insert_entity_links_batch(conn, all_entity_links) logger.debug(f"[9] Batch insert entity links: {time.time() - step_start:.3f}s") # Transaction auto-commits on success commit_start = time.time() logger.debug(f"[10] Commit: {time.time() - commit_start:.3f}s") # 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 logger.debug(f"\n{'='*60}") logger.debug(f"PUT_BATCH_ASYNC COMPLETE: {len(created_unit_ids)} units from {len(contents)} contents in {total_time:.3f}s") logger.debug(f"{'='*60}\n") 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)}") def search( self, agent_id: str, query: str, thinking_budget: int = 50, top_k: int = 10, enable_trace: bool = False, weight_activation: float = 0.30, weight_semantic: float = 0.30, weight_recency: float = 0.25, weight_frequency: float = 0.15, mmr_lambda: float = 0.5, ) -> tuple[List[Dict[str, Any]], Optional[Any]]: """ Search memories using spreading activation (synchronous wrapper). This is a synchronous wrapper around search_async() for convenience. For best performance, use search_async() directly. Args: agent_id: Agent ID to search for query: Search query thinking_budget: How many units to explore (computational budget) top_k: Number of results to return enable_trace: If True, returns detailed SearchTrace object weight_activation: Weight for activation component (default: 0.30) weight_semantic: Weight for semantic similarity component (default: 0.30) weight_recency: Weight for recency component (default: 0.25) weight_frequency: Weight for frequency component (default: 0.15) mmr_lambda: Lambda for MMR diversification (0=max diversity, 1=no diversity, default: 0.5) Returns: Tuple of (results, trace) """ # Run async version synchronously return asyncio.run(self.search_async( agent_id, query, thinking_budget, top_k, enable_trace, weight_activation, weight_semantic, weight_recency, weight_frequency, mmr_lambda )) async def search_async( self, agent_id: str, query: str, thinking_budget: int = 50, top_k: int = 10, enable_trace: bool = False, weight_activation: float = 0.30, weight_semantic: float = 0.30, weight_recency: float = 0.25, weight_frequency: float = 0.15, mmr_lambda: float = 0.5, ) -> tuple[List[Dict[str, Any]], Optional[Any]]: """ Search memories using spreading activation (ASYNC version). This implements the core SEARCH operation: 1. Find entry points (most relevant units via vector search) 2. Spread activation through the graph 3. Weight results by activation + recency + frequency 4. Return top results Args: agent_id: Agent ID to search for query: Search query thinking_budget: How many units to explore (computational budget) top_k: Number of results to return live_tracer: Optional LiveSearchTracer for visualization Returns: List of memory units with their weights, sorted by relevance """ # Initialize tracer if requested from .search_tracer import SearchTracer tracer = SearchTracer(query, thinking_budget, top_k) if enable_trace else None if tracer: tracer.start() pool = await self._get_pool() search_start = time.time() # Buffer logs for clean output in concurrent scenarios search_id = f"{agent_id[:8]}-{int(time.time() * 1000) % 100000}" log_buffer = [] log_buffer.append(f"[SEARCH {search_id}] Query: '{query[:50]}...' (budget={thinking_budget}, top_k={top_k})") try: # Step 1: Generate query embedding (CPU-bound, no DB needed) step_start = time.time() query_embedding = self._generate_embedding(query) step_duration = time.time() - step_start log_buffer.append(f" [1] Generate query embedding: {step_duration:.3f}s") if tracer: tracer.record_query_embedding(query_embedding) tracer.add_phase_metric("generate_query_embedding", step_duration) # Step 2: Find entry points (acquire connection only for this query) step_start = time.time() query_embedding_str = str(query_embedding) # Log connection acquisition conn_acquire_start = time.time() async with pool.acquire() as conn: conn_acquire_time = time.time() - conn_acquire_start if conn_acquire_time > 0.1: # Log if waiting > 100ms log_buffer.append(f" [2.1] Waited {conn_acquire_time:.3f}s for connection (pool busy)") entry_points = await conn.fetch( """ SELECT id, text, context, event_date, access_count, embedding, 1 - (embedding <=> $1::vector) AS similarity FROM memory_units WHERE agent_id = $2 AND embedding IS NOT NULL AND (1 - (embedding <=> $1::vector)) >= 0.5 ORDER BY embedding <=> $1::vector LIMIT 3 """, query_embedding_str, agent_id ) step_duration = time.time() - step_start log_buffer.append(f" [2] Find entry points: {len(entry_points)} found in {step_duration:.3f}s") if tracer: tracer.add_phase_metric("find_entry_points", step_duration, {"count": len(entry_points)}) for rank, ep in enumerate(entry_points, 1): tracer.add_entry_point( node_id=str(ep["id"]), text=ep["text"], similarity=ep["similarity"], rank=rank ) if not entry_points: logger.debug(f"[SEARCH] Complete: 0 results in {time.time() - search_start:.3f}s") if tracer: trace = tracer.finalize([]) return [], trace return [], None # Step 3: Spreading activation with budget (in-memory processing) step_start = time.time() visited = set() results = [] budget_remaining = thinking_budget # Initialize entry points with their actual similarity scores instead of 1.0 # Format: (unit, activation, is_entry, parent_node_id, link_type, link_weight) queue = [(dict(unit), unit["similarity"], True, None, None, None) for unit in entry_points] # Track substep timings calculate_weight_time = 0 query_neighbors_time = 0 process_neighbors_time = 0 # Track which nodes were visited for deferred access count update visited_node_ids = [] # Process nodes in batches for efficient neighbor querying BATCH_SIZE = 50 nodes_to_process = [] # (unit, activation, is_entry_point, parent_node_id, link_type, link_weight) while queue and budget_remaining > 0: # Collect a batch of nodes to process (in-memory, no DB) while queue and len(nodes_to_process) < BATCH_SIZE and budget_remaining > 0: current_unit, activation, is_entry_point, parent_node_id, link_type, link_weight = queue.pop(0) unit_id = str(current_unit["id"]) if unit_id not in visited: visited.add(unit_id) budget_remaining -= 1 nodes_to_process.append((current_unit, activation, is_entry_point, parent_node_id, link_type, link_weight)) visited_node_ids.append(unit_id) # Track for deferred update elif tracer: # Node already visited - prune tracer.prune_node(unit_id, "already_visited", activation) if not nodes_to_process: break # Acquire connection ONLY for neighbor queries (defer access count updates) node_ids = [str(node[0]["id"]) for node in nodes_to_process] # Log connection acquisition for batch queries batch_conn_start = time.time() async with pool.acquire() as conn: batch_conn_acquire = time.time() - batch_conn_start if batch_conn_acquire > 0.1: # Log if waiting > 100ms log_buffer.append(f" [3.3.1] Waited {batch_conn_acquire:.3f}s for connection (pool busy) - batch size: {len(node_ids)}") # Query neighbors for ALL nodes in batch at once (without embeddings for speed) # Convert string UUIDs to UUID type for faster matching substep_start = time.time() uuid_array = [uuid.UUID(nid) for nid in node_ids] all_neighbors = await conn.fetch( """ SELECT ml.from_unit_id, ml.to_unit_id, ml.weight, ml.link_type, ml.entity_id, mu.text, mu.context, mu.event_date, mu.access_count, mu.id as neighbor_id FROM memory_links ml JOIN memory_units mu ON ml.to_unit_id = mu.id WHERE ml.from_unit_id = ANY($1::uuid[]) AND ml.weight >= 0.1 ORDER BY ml.from_unit_id, ml.weight DESC """, uuid_array ) neighbor_query_time = time.time() - substep_start if neighbor_query_time > 1.0: # Log slow neighbor queries log_buffer.append(f" [3.3.3] Slow NEIGHBOR query: {neighbor_query_time:.3f}s for {len(node_ids)} nodes → {len(all_neighbors)} neighbors") query_neighbors_time += neighbor_query_time # Fetch embeddings for current batch nodes (needed for weight calculation) substep_start = time.time() embeddings = await conn.fetch( "SELECT id, embedding FROM memory_units WHERE id = ANY($1::uuid[])", uuid_array ) embedding_map = {str(row["id"]): row["embedding"] for row in embeddings} fetch_embeddings_time = time.time() - substep_start if fetch_embeddings_time > 0.5: log_buffer.append(f" [3.3.4] Slow EMBEDDING fetch: {fetch_embeddings_time:.3f}s for {len(node_ids)} nodes") query_neighbors_time += fetch_embeddings_time # Group neighbors by from_unit_id (in-memory, no DB) substep_start = time.time() neighbors_by_node = {} for neighbor in all_neighbors: from_id = str(neighbor["from_unit_id"]) if from_id not in neighbors_by_node: neighbors_by_node[from_id] = [] neighbors_by_node[from_id].append(neighbor) # Process each node in the batch (CPU-bound, no DB) for current_unit, activation, is_entry_point, parent_node_id, parent_link_type, parent_link_weight in nodes_to_process: unit_id = str(current_unit["id"]) # Calculate combined weight event_date = current_unit["event_date"] days_since = (utcnow() - event_date).total_seconds() / 86400 recency_weight = calculate_recency_weight(days_since) frequency_weight = calculate_frequency_weight(current_unit.get("access_count", 0)) # Normalize frequency to [0, 1] range frequency_normalized = (frequency_weight - 1.0) / 1.0 # Calculate semantic similarity between query and this memory # Get embedding from the map we fetched memory_embedding = embedding_map.get(unit_id) if memory_embedding is not None: # Convert embedding to list of floats if it's a string or other type if isinstance(memory_embedding, str): import json memory_embedding = json.loads(memory_embedding) elif not isinstance(memory_embedding, (list, np.ndarray)): # If it's some other type, try to convert it memory_embedding = list(memory_embedding) # Cosine similarity = 1 - cosine distance query_vec = np.array(query_embedding, dtype=np.float64) memory_vec = np.array(memory_embedding, dtype=np.float64) # Cosine similarity dot_product = np.dot(query_vec, memory_vec) norm_query = np.linalg.norm(query_vec) norm_memory = np.linalg.norm(memory_vec) semantic_similarity = dot_product / (norm_query * norm_memory) if norm_query > 0 and norm_memory > 0 else 0.0 else: semantic_similarity = 0.0 # Combined weight using configurable parameters final_weight = ( weight_activation * activation + weight_semantic * semantic_similarity + weight_recency * recency_weight + weight_frequency * frequency_normalized ) # Notify tracer if tracer: tracer.visit_node( node_id=unit_id, text=current_unit["text"], context=current_unit.get("context", ""), event_date=event_date, access_count=current_unit.get("access_count", 0), is_entry_point=is_entry_point, parent_node_id=parent_node_id, link_type=parent_link_type, link_weight=parent_link_weight, activation=activation, semantic_similarity=semantic_similarity, recency=recency_weight, frequency=frequency_normalized, final_weight=final_weight, ) results.append({ "id": unit_id, "text": current_unit["text"], "context": current_unit.get("context", ""), "event_date": event_date.isoformat(), "weight": final_weight, "activation": activation, "semantic_similarity": semantic_similarity, "recency": recency_weight, "frequency": frequency_weight, "embedding": memory_embedding, # Store for MMR }) # Spread to neighbors (from batch query results) neighbors = neighbors_by_node.get(unit_id, []) # Group neighbors by to_unit_id to handle multiple connections neighbors_grouped = {} for neighbor in neighbors: neighbor_id = str(neighbor["to_unit_id"]) if neighbor_id not in neighbors_grouped: neighbors_grouped[neighbor_id] = [] neighbors_grouped[neighbor_id].append(neighbor) # Process each unique neighbor (aggregating multiple links) for neighbor_id, neighbor_links in neighbors_grouped.items(): if neighbor_id in visited: continue # Sort links by weight descending to identify primary link neighbor_links_sorted = sorted(neighbor_links, key=lambda x: x["weight"], reverse=True) primary_link = neighbor_links_sorted[0] # Aggregate link weights: max + 30% bonus for additional links max_weight = primary_link["weight"] bonus_weight = sum(link["weight"] for link in neighbor_links_sorted[1:]) * 0.3 combined_weight = max_weight + bonus_weight # Calculate new activation using combined weight new_activation = activation * combined_weight * 0.8 # 0.8 = decay factor # Use primary link metadata for queue and trace primary_link_type = primary_link["link_type"] primary_entity_id = str(primary_link["entity_id"]) if primary_link["entity_id"] else None if new_activation > 0.1: queue.append(({ "id": primary_link["to_unit_id"], "text": primary_link["text"], "context": primary_link.get("context", ""), "event_date": primary_link["event_date"], "access_count": primary_link["access_count"], }, new_activation, False, unit_id, primary_link_type, combined_weight)) # parent_id, link_type, combined_weight # Record all links in trace (primary + additional) if tracer: # Add primary link with combined activation tracer.add_neighbor_link( from_node_id=unit_id, to_node_id=neighbor_id, link_type=primary_link_type, link_weight=combined_weight, entity_id=primary_entity_id, new_activation=new_activation, followed=True ) # Add additional links as supplementary (if multiple connections exist) for additional_link in neighbor_links_sorted[1:]: additional_link_type = additional_link["link_type"] additional_entity_id = str(additional_link["entity_id"]) if additional_link["entity_id"] else None tracer.add_neighbor_link( from_node_id=unit_id, to_node_id=neighbor_id, link_type=additional_link_type, link_weight=additional_link["weight"], entity_id=additional_entity_id, new_activation=None, # Don't show activation for supplementary links followed=True, is_supplementary=True # Mark as supplementary link ) elif tracer: # Record pruned link tracer.add_neighbor_link( from_node_id=unit_id, to_node_id=neighbor_id, link_type=primary_link_type, link_weight=combined_weight, entity_id=primary_entity_id, new_activation=new_activation, followed=False, prune_reason="activation_too_low" ) calculate_weight_time += time.time() - substep_start process_neighbors_time += time.time() - substep_start # Clear batch for next iteration nodes_to_process = [] spreading_activation_time = time.time() - step_start num_batches = (len(visited) + BATCH_SIZE - 1) // BATCH_SIZE # Ceiling division log_buffer.append(f" [3] Spreading activation: {len(visited)} nodes visited in {spreading_activation_time:.3f}s") log_buffer.append(f" [3.1] Calculate weights: {calculate_weight_time:.3f}s") log_buffer.append(f" [3.2] Query neighbors: {query_neighbors_time:.3f}s ({num_batches} batched queries)") log_buffer.append(f" [3.3] Process neighbors: {process_neighbors_time:.3f}s") if tracer: tracer.add_phase_metric("spreading_activation", spreading_activation_time, { "nodes_visited": len(visited), "num_batches": num_batches }) # Step 4: Queue access count updates (background worker will process them) if visited_node_ids: await self._access_count_queue.put(visited_node_ids) log_buffer.append(f" [4] Queued access count updates for {len(visited_node_ids)} nodes") # Step 5: Sort by final weight and apply MMR for diversity step_start = time.time() results.sort(key=lambda x: x["weight"], reverse=True) # Apply MMR (Maximal Marginal Relevance) for diversity if lambda < 1.0 if mmr_lambda < 1.0 and len(results) > top_k: top_results = self._apply_mmr(results, top_k, mmr_lambda, log_buffer) log_buffer.append(f" [5] MMR diversification (λ={mmr_lambda}): {time.time() - step_start:.3f}s") else: top_results = results[:top_k] # Add original rank and remove embeddings from results for idx, result in enumerate(top_results): result["original_rank"] = idx + 1 result["mmr_score"] = None result["mmr_relevance"] = None result["mmr_max_similarity"] = None result["mmr_diversified"] = False result.pop("embedding", None) log_buffer.append(f" [5] Sort and return top {top_k} (no MMR): {time.time() - step_start:.3f}s") total_time = time.time() - search_start log_buffer.append(f"[SEARCH {search_id}] Complete: {len(top_results)} results in {total_time:.3f}s") # Log all buffered logs at once logger.info("\n" + "\n".join(log_buffer)) # Finalize trace if enabled if tracer: trace = tracer.finalize(top_results) return top_results, trace return top_results, None except Exception as e: log_buffer.append(f"[SEARCH {search_id}] ERROR after {time.time() - search_start:.3f}s: {str(e)}") logger.error("\n" + "\n".join(log_buffer)) raise Exception(f"Failed to search memories: {str(e)}") def _apply_mmr( self, results: List[Dict[str, Any]], top_k: int, mmr_lambda: float, log_buffer: List[str] ) -> List[Dict[str, Any]]: """ Apply Maximal Marginal Relevance (MMR) to diversify search results. MMR balances relevance with diversity by selecting results that are: 1. Relevant to the query (high score) 2. Different from already selected results (low similarity) Formula: MMR = λ * relevance - (1-λ) * max_similarity_to_selected Args: results: Sorted list of all results with embeddings top_k: Number of results to select mmr_lambda: Balance parameter (0=max diversity, 1=max relevance) log_buffer: Logging buffer Returns: Diversified list of top_k results """ if not results or top_k <= 0: return [] # Normalize weights to [0, 1] for fair comparison with similarity max_weight = max(r["weight"] for r in results) min_weight = min(r["weight"] for r in results) weight_range = max_weight - min_weight if max_weight > min_weight else 1.0 # Pre-compute normalized relevance scores for all results for idx, result in enumerate(results): result["original_rank"] = idx + 1 result["normalized_relevance"] = (result["weight"] - min_weight) / weight_range # Extract embeddings as a numpy array for vectorized operations # Shape: (num_results, embedding_dim) embeddings_list = [] valid_indices = [] for idx, result in enumerate(results): if result.get("embedding") is not None: embeddings_list.append(result["embedding"]) valid_indices.append(idx) if not embeddings_list: # No embeddings available, just return top-k by relevance return results[:top_k] # Stack embeddings into a matrix (num_results, embedding_dim) embeddings_matrix = np.array(embeddings_list, dtype=np.float32) # Normalize embeddings for faster cosine similarity (just dot product after normalization) norms = np.linalg.norm(embeddings_matrix, axis=1, keepdims=True) norms[norms == 0] = 1.0 # Avoid division by zero embeddings_matrix = embeddings_matrix / norms selected_indices = [] remaining_indices = list(range(len(results))) diversified_count = 0 for selection_round in range(min(top_k, len(results))): if not remaining_indices: break best_mmr_score = float('-inf') best_remaining_idx = 0 # Vectorized computation for all remaining candidates for remaining_idx, candidate_idx in enumerate(remaining_indices): candidate = results[candidate_idx] normalized_relevance = candidate["normalized_relevance"] # Calculate max similarity to selected results max_similarity = 0.0 if selected_indices and candidate_idx in valid_indices: # Find position in embeddings_matrix embedding_idx = valid_indices.index(candidate_idx) candidate_embedding = embeddings_matrix[embedding_idx] # Vectorized similarity calculation with all selected embeddings if selected_indices: selected_embedding_indices = [valid_indices.index(idx) for idx in selected_indices if idx in valid_indices] if selected_embedding_indices: selected_embeddings = embeddings_matrix[selected_embedding_indices] # Compute cosine similarities in one operation (already normalized, so just dot product) similarities = np.dot(selected_embeddings, candidate_embedding) max_similarity = float(np.max(similarities)) # MMR score: balance relevance and diversity mmr_score = mmr_lambda * normalized_relevance - (1 - mmr_lambda) * max_similarity if mmr_score > best_mmr_score: best_mmr_score = mmr_score best_remaining_idx = remaining_idx best_max_similarity = max_similarity # Select the best candidate best_candidate_idx = remaining_indices.pop(best_remaining_idx) best_candidate = results[best_candidate_idx] # Store MMR metadata best_candidate["mmr_score"] = best_mmr_score best_candidate["mmr_relevance"] = best_candidate["normalized_relevance"] best_candidate["mmr_max_similarity"] = best_max_similarity best_candidate["mmr_diversified"] = best_remaining_idx > 0 selected_indices.append(best_candidate_idx) if best_remaining_idx > 0: diversified_count += 1 log_buffer.append(f" MMR: Selected {len(selected_indices)} results, {diversified_count} diversified picks") # Return selected results in order selected_results = [results[idx] for idx in selected_indices] # Remove embeddings from final results (not needed in response) for result in selected_results: result.pop("embedding", None) result.pop("normalized_relevance", None) # Clean up temp field return selected_results async def get_document(self, document_id: str, agent_id: str) -> Optional[Dict[str, Any]]: """ Retrieve document metadata and statistics. Args: document_id: Document ID to retrieve agent_id: Agent ID that owns the document Returns: Dictionary with document info or None if not found """ pool = await self._get_pool() async with pool.acquire() as conn: doc = await conn.fetchrow( """ SELECT d.id, d.agent_id, d.original_text, d.content_hash, d.metadata, d.created_at, d.updated_at, COUNT(mu.id) as unit_count FROM documents d LEFT JOIN memory_units mu ON mu.document_id = d.id WHERE d.id = $1 AND d.agent_id = $2 GROUP BY d.id, d.agent_id, d.original_text, d.content_hash, d.metadata, d.created_at, d.updated_at """, document_id, agent_id ) if not doc: return None import json return { "id": doc["id"], "agent_id": doc["agent_id"], "original_text": doc["original_text"], "content_hash": doc["content_hash"], "metadata": json.loads(doc["metadata"]) if doc["metadata"] else {}, "unit_count": doc["unit_count"], "created_at": doc["created_at"], "updated_at": doc["updated_at"] } async def delete_document(self, document_id: str, agent_id: str) -> Dict[str, int]: """ Delete a document and all its associated memory units and links. Args: document_id: Document ID to delete agent_id: Agent ID that owns the document Returns: Dictionary with counts of deleted items """ pool = await self._get_pool() async with pool.acquire() as conn: async with conn.transaction(): # Count units before deletion units_count = await conn.fetchval( "SELECT COUNT(*) FROM memory_units WHERE document_id = $1", document_id ) # Delete document (cascades to memory_units and all their links) deleted = await conn.fetchval( "DELETE FROM documents WHERE id = $1 AND agent_id = $2 RETURNING id", document_id, agent_id ) return { "document_deleted": 1 if deleted else 0, "memory_units_deleted": units_count if deleted else 0 } async def delete_agent(self, agent_id: str) -> Dict[str, int]: """ Delete all data for a specific agent (multi-tenant cleanup). This is much more efficient than dropping all tables and allows multiple agents to coexist in the same database. Deletes (with CASCADE): - All memory units for this agent - All entities for this agent - All associated links, unit-entity associations, and co-occurrences Args: agent_id: Agent ID to delete Returns: Dictionary with counts of deleted items """ pool = await self._get_pool() async with pool.acquire() as conn: async with conn.transaction(): try: # Count before deletion for reporting units_count = await conn.fetchval("SELECT COUNT(*) FROM memory_units WHERE agent_id = $1", agent_id) entities_count = await conn.fetchval("SELECT COUNT(*) FROM entities WHERE agent_id = $1", agent_id) # Delete memory units (cascades to unit_entities, memory_links) await conn.execute("DELETE FROM memory_units WHERE agent_id = $1", agent_id) # Delete entities (cascades to unit_entities, entity_cooccurrences, memory_links with entity_id) await conn.execute("DELETE FROM entities WHERE agent_id = $1", agent_id) return { "memory_units_deleted": units_count, "entities_deleted": entities_count } except Exception as e: raise Exception(f"Failed to delete agent data: {str(e)}") async def _extract_entities_batch_optimized( self, conn, agent_id: str, unit_ids: List[str], sentences: List[str], context: str, fact_dates: List, llm_entities: List[List[Dict]], # NEW: Entities from LLM ) -> List[tuple]: """ Process LLM-extracted entities for ALL facts in batch. Uses entities provided by the LLM (no spaCy needed), then resolves and links them in bulk. Returns list of tuples for batch insertion: (from_unit_id, to_unit_id, link_type, weight, entity_id) """ try: # Step 1: Convert LLM entities to the format expected by entity resolver substep_start = time.time() all_entities = [] for entity_list in llm_entities: # Convert List[Entity] or List[dict] to List[Dict] format formatted_entities = [] for ent in entity_list: # Handle both Entity objects and dicts if hasattr(ent, 'text'): formatted_entities.append({'text': ent.text, 'type': ent.type}) elif isinstance(ent, dict): formatted_entities.append({'text': ent.get('text', ''), 'type': ent.get('type', 'CONCEPT')}) all_entities.append(formatted_entities) total_entities = sum(len(ents) for ents in all_entities) logger.debug(f" [6.1] Process LLM entities: {total_entities} entities from {len(sentences)} facts in {time.time() - substep_start:.3f}s") # Step 2: Resolve entities in BATCH (much faster!) substep_start = time.time() step_6_2_start = time.time() # [6.2.1] Prepare all entities for batch resolution substep_6_2_1_start = time.time() all_entities_flat = [] entity_to_unit = [] # Maps flat index to (unit_id, local_index) for unit_id, entities, fact_date in zip(unit_ids, all_entities, fact_dates): if not entities: continue for local_idx, entity in enumerate(entities): all_entities_flat.append({ 'text': entity['text'], 'type': entity['type'], 'nearby_entities': entities, }) entity_to_unit.append((unit_id, local_idx, fact_date)) logger.debug(f" [6.2.1] Prepare entities: {len(all_entities_flat)} entities in {time.time() - substep_6_2_1_start:.3f}s") # Resolve ALL entities in one batch call if all_entities_flat: # [6.2.2] Batch resolve entities substep_6_2_2_start = time.time() # Group by date for batch resolution (most will have same date) entities_by_date = {} for idx, (unit_id, local_idx, fact_date) in enumerate(entity_to_unit): date_key = fact_date if date_key not in entities_by_date: entities_by_date[date_key] = [] entities_by_date[date_key].append((idx, all_entities_flat[idx])) # Resolve each date group in batch resolved_entity_ids = [None] * len(all_entities_flat) for fact_date, entities_group in entities_by_date.items(): indices = [idx for idx, _ in entities_group] entities_data = [entity_data for _, entity_data in entities_group] batch_resolved = await self.entity_resolver.resolve_entities_batch( agent_id=agent_id, entities_data=entities_data, context=context, unit_event_date=fact_date, conn=conn ) for idx, entity_id in zip(indices, batch_resolved): resolved_entity_ids[idx] = entity_id logger.debug(f" [6.2.2] Resolve entities: {len(all_entities_flat)} entities in {time.time() - substep_6_2_2_start:.3f}s") # [6.2.3] Create unit-entity links in BATCH substep_6_2_3_start = time.time() # Map resolved entities back to units and collect all (unit, entity) pairs unit_to_entity_ids = {} unit_entity_pairs = [] for idx, (unit_id, local_idx, fact_date) in enumerate(entity_to_unit): if unit_id not in unit_to_entity_ids: unit_to_entity_ids[unit_id] = [] entity_id = resolved_entity_ids[idx] unit_to_entity_ids[unit_id].append(entity_id) unit_entity_pairs.append((unit_id, entity_id)) # Batch insert all unit-entity links (MUCH faster!) await self.entity_resolver.link_units_to_entities_batch(unit_entity_pairs, conn=conn) logger.debug(f" [6.2.3] Create unit-entity links (batched): {len(unit_entity_pairs)} links in {time.time() - substep_6_2_3_start:.3f}s") logger.debug(f" [6.2] Entity resolution (batched): {len(all_entities_flat)} entities resolved in {time.time() - step_6_2_start:.3f}s") else: unit_to_entity_ids = {} logger.debug(f" [6.2] Entity resolution (batched): 0 entities in {time.time() - step_6_2_start:.3f}s") # Step 3: Create entity links between units that share entities substep_start = time.time() # Collect all unique entity IDs all_entity_ids = set() for entity_ids in unit_to_entity_ids.values(): all_entity_ids.update(entity_ids) # For each entity, find all units that reference it (one query per entity) entity_to_units = {} for entity_id in all_entity_ids: rows = await conn.fetch( """ SELECT unit_id FROM unit_entities WHERE entity_id = $1 """, entity_id ) entity_to_units[entity_id] = [row['unit_id'] for row in rows] # Create bidirectional links between units that share entities links = [] for entity_id, units_with_entity in entity_to_units.items(): # For each pair of units with this entity, create bidirectional links for i, unit_id_1 in enumerate(units_with_entity): for unit_id_2 in units_with_entity[i+1:]: # Bidirectional links links.append((unit_id_1, unit_id_2, 'entity', 1.0, entity_id)) links.append((unit_id_2, unit_id_1, 'entity', 1.0, entity_id)) logger.debug(f" [6.3] Entity link creation: {len(links)} links for {len(all_entity_ids)} unique entities in {time.time() - substep_start:.3f}s") return links except Exception as e: logger.error(f" Failed to extract entities in batch: {str(e)}") import traceback traceback.print_exc() # Re-raise to trigger rollback at put_async level raise async def _create_temporal_links_batch_per_fact( self, conn, agent_id: str, unit_ids: List[str], time_window_hours: int = 24, ): """ Create temporal links for multiple units, each with their own event_date. Queries the event_date for each unit from the database and creates temporal links based on individual dates (supports per-fact dating). """ if not unit_ids: return try: # Get the event_date for each new unit rows = await conn.fetch( """ SELECT id, event_date FROM memory_units WHERE id::text = ANY($1) """, unit_ids ) new_units = {str(row['id']): row['event_date'] for row in rows} # Create links based on each unit's individual event_date links = [] for unit_id, unit_event_date in new_units.items(): # Find units within the time window of THIS specific unit recent_units = await conn.fetch( """ SELECT id, event_date FROM memory_units WHERE agent_id = $1 AND id != $2 AND event_date BETWEEN $3 AND $4 ORDER BY event_date DESC LIMIT 10 """, agent_id, unit_id, unit_event_date - timedelta(hours=time_window_hours), unit_event_date + timedelta(hours=time_window_hours) ) for recent_row in recent_units: recent_id = recent_row['id'] recent_event_date = recent_row['event_date'] # Calculate temporal proximity weight time_diff_hours = abs((unit_event_date - recent_event_date).total_seconds() / 3600) weight = max(0.3, 1.0 - (time_diff_hours / time_window_hours)) links.append((unit_id, str(recent_id), 'temporal', weight, None)) if links: await conn.executemany( """ INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id) VALUES ($1, $2, $3, $4, $5) ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING """, links ) except Exception as e: logger.error(f" Failed to create temporal links: {str(e)}") import traceback traceback.print_exc() # Re-raise to trigger rollback at put_async level raise async def _create_semantic_links_batch( self, conn, agent_id: str, unit_ids: List[str], embeddings: List[List[float]], top_k: int = 5, threshold: float = 0.7, ): """ Create semantic links for multiple units efficiently. For each unit, finds similar units and creates links. """ if not unit_ids or not embeddings: return try: all_links = [] for unit_id, embedding in zip(unit_ids, embeddings): # Find similar units using vector similarity # Convert embedding to string for asyncpg embedding_str = str(embedding) similar_units = await conn.fetch( """ SELECT id, 1 - (embedding <=> $1::vector) AS similarity FROM memory_units WHERE agent_id = $2 AND id != $3 AND embedding IS NOT NULL AND (1 - (embedding <=> $1::vector)) >= $4 ORDER BY embedding <=> $1::vector LIMIT $5 """, embedding_str, agent_id, unit_id, threshold, top_k ) for row in similar_units: similar_id = row['id'] similarity = row['similarity'] all_links.append((unit_id, str(similar_id), 'semantic', float(similarity), None)) if all_links: await conn.executemany( """ INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id) VALUES ($1, $2, $3, $4, $5) ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING """, all_links ) except Exception as e: logger.error(f" Failed to create semantic links: {str(e)}") import traceback traceback.print_exc() # Re-raise to trigger rollback at put_async level raise async def _insert_entity_links_batch(self, conn, links: List[tuple]): """Insert all entity links in a single batch.""" if not links: return try: await conn.executemany( """ INSERT INTO memory_links (from_unit_id, to_unit_id, link_type, weight, entity_id) VALUES ($1, $2, $3, $4, $5) ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING """, links ) except Exception as e: logger.warning(f" Failed to insert entity links: {str(e)}")