""" Entity extraction and resolution for memory system. Uses spaCy for entity extraction and implements resolution logic to disambiguate entities across memory units. """ import asyncpg from typing import List, Dict, Optional, Set from difflib import SequenceMatcher from datetime import datetime, timezone # Load spaCy model (singleton) _nlp = None class EntityResolver: """ Resolves entities to canonical IDs with disambiguation. """ def __init__(self, pool: asyncpg.Pool): """ Initialize entity resolver. Args: pool: asyncpg connection pool """ self.pool = pool async def resolve_entities_batch( self, agent_id: str, entities_data: List[Dict], context: str, unit_event_date, conn=None, ) -> List[str]: """ Resolve multiple entities in batch (MUCH faster than sequential). Groups entities by type, queries candidates in bulk, and resolves all entities with minimal DB queries. Args: agent_id: Agent ID entities_data: List of dicts with 'text', 'type', 'nearby_entities' context: Context where entities appear unit_event_date: When this unit was created conn: Optional connection to use (if None, acquires from pool) Returns: List of entity IDs in same order as input """ if not entities_data: return [] if conn is None: async with self.pool.acquire() as conn: return await self._resolve_entities_batch_impl(conn, agent_id, entities_data, context, unit_event_date) else: return await self._resolve_entities_batch_impl(conn, agent_id, entities_data, context, unit_event_date) async def _resolve_entities_batch_impl(self, conn, agent_id: str, entities_data: List[Dict], context: str, unit_event_date) -> List[str]: import time start = time.time() # Group entities by type for efficient querying entities_by_type = {} for idx, entity_data in enumerate(entities_data): entity_type = entity_data['type'] if entity_type not in entities_by_type: entities_by_type[entity_type] = [] entities_by_type[entity_type].append((idx, entity_data)) # Query ALL candidates for each type in batch all_candidates = {} # Maps (entity_type, entity_text) -> list of candidates for entity_type, entities_list in entities_by_type.items(): # Extract unique entity texts for this type entity_texts = list(set(e[1]['text'] for e in entities_list)) # Query candidates for all texts at once type_candidates = await conn.fetch( """ SELECT canonical_name, id, metadata, last_seen, mention_count FROM entities WHERE agent_id = $1 AND entity_type = $2 """, agent_id, entity_type ) # Filter candidates in memory (faster than complex SQL for small datasets) for entity_text in entity_texts: matching = [] entity_text_lower = entity_text.lower() for row in type_candidates: canonical_name = row['canonical_name'] ent_id = row['id'] metadata = row['metadata'] last_seen = row['last_seen'] mention_count = row['mention_count'] canonical_lower = canonical_name.lower() # Same matching logic as before if (entity_text_lower == canonical_lower or entity_text_lower in canonical_lower or canonical_lower in entity_text_lower): matching.append((ent_id, canonical_name, metadata, last_seen, mention_count)) all_candidates[(entity_type, entity_text)] = matching # Resolve each entity using pre-fetched candidates entity_ids = [None] * len(entities_data) entities_to_update = [] # (entity_id, unit_event_date) entities_to_create = [] # (idx, entity_data) for idx, entity_data in enumerate(entities_data): entity_text = entity_data['text'] entity_type = entity_data['type'] nearby_entities = entity_data.get('nearby_entities', []) candidates = all_candidates.get((entity_type, entity_text), []) if not candidates: # Will create new entity entities_to_create.append((idx, entity_data)) continue # Score candidates (same logic as before but with pre-fetched data) best_candidate = None best_score = 0.0 best_name_similarity = 0.0 nearby_entity_set = {e['text'].lower() for e in nearby_entities if e['text'] != entity_text} for candidate_id, canonical_name, metadata, last_seen, mention_count in candidates: score = 0.0 # Name similarity name_similarity = SequenceMatcher( None, entity_text.lower(), canonical_name.lower() ).ratio() score += name_similarity * 0.5 # Temporal proximity if last_seen: days_diff = abs((unit_event_date - last_seen).total_seconds() / 86400) if days_diff < 7: temporal_score = max(0, 1.0 - (days_diff / 7)) score += temporal_score * 0.2 if score > best_score: best_score = score best_candidate = candidate_id best_name_similarity = name_similarity # Apply threshold threshold = 0.4 if entity_type == 'PERSON' and best_name_similarity >= 0.95 else 0.6 if best_score > threshold: entity_ids[idx] = best_candidate entities_to_update.append((best_candidate, unit_event_date)) else: entities_to_create.append((idx, entity_data)) # Batch update existing entities if entities_to_update: await conn.executemany( """ UPDATE entities SET mention_count = mention_count + 1, last_seen = $2 WHERE id = $1::uuid """, entities_to_update ) # Batch create new entities using multi-row VALUES if entities_to_create: import logging logger = logging.getLogger(__name__) create_start = time.time() # Build multi-row VALUES statement # VALUES ($1, $2, ...), ($N+1, $N+2, ...), ... values_clauses = [] params = [] param_idx = 1 for idx, entity_data in entities_to_create: values_clauses.append(f"(${param_idx}, ${param_idx+1}, ${param_idx+2}, ${param_idx+3}, ${param_idx+4}, ${param_idx+5})") params.extend([ agent_id, entity_data['text'], entity_data['type'], unit_event_date, unit_event_date, 1 ]) param_idx += 6 # Single INSERT with multiple VALUES rows query = f""" INSERT INTO entities (agent_id, canonical_name, entity_type, first_seen, last_seen, mention_count) VALUES {', '.join(values_clauses)} RETURNING id """ created_rows = await conn.fetch(query, *params) # Map created IDs back to original indices for i, (idx, entity_data) in enumerate(entities_to_create): entity_ids[idx] = created_rows[i]['id'] logger.info(f" [6.2.2.X] Batch created {len(entities_to_create)} new entities in {time.time() - create_start:.3f}s") return entity_ids async def resolve_entity( self, agent_id: str, entity_text: str, entity_type: str, context: str, nearby_entities: List[Dict], unit_event_date, ) -> str: """ Resolve an entity to a canonical entity ID. Args: agent_id: Agent ID (entities are scoped to agents) entity_text: Entity text ("Alice", "Google", etc.) entity_type: Entity type (PERSON, ORG, etc.) context: Context where entity appears nearby_entities: Other entities in the same unit unit_event_date: When this unit was created Returns: Entity ID (creates new entity if needed) """ async with self.pool.acquire() as conn: # Find candidate entities with same type and similar name candidates = await conn.fetch( """ SELECT id, canonical_name, metadata, last_seen FROM entities WHERE agent_id = $1 AND entity_type = $2 AND ( canonical_name ILIKE $3 OR canonical_name ILIKE $4 OR $3 ILIKE canonical_name || '%%' ) ORDER BY mention_count DESC """, agent_id, entity_type, entity_text, f"%{entity_text}%" ) if not candidates: # New entity - create it return await self._create_entity( conn, agent_id, entity_text, entity_type, unit_event_date ) # Score candidates based on: # 1. Name similarity # 2. Context overlap (TODO: could use embeddings) # 3. Co-occurring entities # 4. Temporal proximity best_candidate = None best_score = 0.0 best_name_similarity = 0.0 nearby_entity_set = {e['text'].lower() for e in nearby_entities if e['text'] != entity_text} for row in candidates: candidate_id = row['id'] canonical_name = row['canonical_name'] metadata = row['metadata'] last_seen = row['last_seen'] score = 0.0 # 1. Name similarity (0-1) name_similarity = SequenceMatcher( None, entity_text.lower(), canonical_name.lower() ).ratio() score += name_similarity * 0.5 # 2. Co-occurring entities (0-0.5) # Get entities that co-occurred with this candidate before # Use the materialized co-occurrence cache for fast lookup co_entity_rows = await conn.fetch( """ SELECT e.canonical_name, ec.cooccurrence_count FROM entity_cooccurrences ec JOIN entities e ON ( CASE WHEN ec.entity_id_1 = $1 THEN ec.entity_id_2 WHEN ec.entity_id_2 = $1 THEN ec.entity_id_1 END = e.id ) WHERE ec.entity_id_1 = $1 OR ec.entity_id_2 = $1 """, candidate_id ) co_entities = {r['canonical_name'].lower() for r in co_entity_rows} # Check overlap with nearby entities overlap = len(nearby_entity_set & co_entities) if nearby_entity_set: co_entity_score = overlap / len(nearby_entity_set) score += co_entity_score * 0.3 # 3. Temporal proximity (0-0.2) if last_seen: days_diff = abs((unit_event_date - last_seen).total_seconds() / 86400) if days_diff < 7: # Within a week temporal_score = max(0, 1.0 - (days_diff / 7)) score += temporal_score * 0.2 if score > best_score: best_score = score best_candidate = candidate_id best_name_similarity = name_similarity # Threshold for considering it the same entity # For PERSON entities with exact name match, use lower threshold threshold = 0.4 if entity_type == 'PERSON' and best_name_similarity >= 0.95 else 0.6 if best_score > threshold: # Update entity await conn.execute( """ UPDATE entities SET mention_count = mention_count + 1, last_seen = $1 WHERE id = $2 """, unit_event_date, best_candidate ) return best_candidate else: # Not confident - create new entity return await self._create_entity( conn, agent_id, entity_text, entity_type, unit_event_date ) async def _create_entity( self, conn, agent_id: str, entity_text: str, entity_type: str, event_date, ) -> str: """ Create a new entity. Args: conn: Database connection agent_id: Agent ID entity_text: Entity text entity_type: Entity type event_date: When first seen Returns: Entity ID """ entity_id = await conn.fetchval( """ INSERT INTO entities (agent_id, canonical_name, entity_type, first_seen, last_seen, mention_count) VALUES ($1, $2, $3, $4, $5, 1) RETURNING id """, agent_id, entity_text, entity_type, event_date, event_date ) return entity_id async def link_unit_to_entity(self, unit_id: str, entity_id: str): """ Link a memory unit to an entity. Also updates co-occurrence cache with other entities in the same unit. Args: unit_id: Memory unit ID entity_id: Entity ID """ async with self.pool.acquire() as conn: # Insert unit-entity link await conn.execute( """ INSERT INTO unit_entities (unit_id, entity_id) VALUES ($1, $2) ON CONFLICT DO NOTHING """, unit_id, entity_id ) # Update co-occurrence cache: find other entities in this unit rows = await conn.fetch( """ SELECT entity_id FROM unit_entities WHERE unit_id = $1 AND entity_id != $2 """, unit_id, entity_id ) other_entities = [row['entity_id'] for row in rows] # Update co-occurrences for each pair for other_entity_id in other_entities: await self._update_cooccurrence(conn, entity_id, other_entity_id) async def _update_cooccurrence(self, conn, entity_id_1: str, entity_id_2: str): """ Update the co-occurrence cache for two entities. Uses CHECK constraint ordering (entity_id_1 < entity_id_2) to avoid duplicates. Args: conn: Database connection entity_id_1: First entity ID entity_id_2: Second entity ID """ # Ensure consistent ordering (smaller UUID first) if entity_id_1 > entity_id_2: entity_id_1, entity_id_2 = entity_id_2, entity_id_1 await conn.execute( """ INSERT INTO entity_cooccurrences (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred) VALUES ($1, $2, 1, NOW()) ON CONFLICT (entity_id_1, entity_id_2) DO UPDATE SET cooccurrence_count = entity_cooccurrences.cooccurrence_count + 1, last_cooccurred = NOW() """, entity_id_1, entity_id_2 ) async def link_units_to_entities_batch(self, unit_entity_pairs: List[tuple[str, str]], conn=None): """ Link multiple memory units to entities in batch (MUCH faster than sequential). Also updates co-occurrence cache for entities that appear in the same unit. Args: unit_entity_pairs: List of (unit_id, entity_id) tuples conn: Optional connection to use (if None, acquires from pool) """ if not unit_entity_pairs: return if conn is None: async with self.pool.acquire() as conn: return await self._link_units_to_entities_batch_impl(conn, unit_entity_pairs) else: return await self._link_units_to_entities_batch_impl(conn, unit_entity_pairs) async def _link_units_to_entities_batch_impl(self, conn, unit_entity_pairs: List[tuple[str, str]]): # Batch insert all unit-entity links await conn.executemany( """ INSERT INTO unit_entities (unit_id, entity_id) VALUES ($1, $2) ON CONFLICT DO NOTHING """, unit_entity_pairs ) # Build map of unit -> entities for co-occurrence calculation # Use sets to avoid duplicate entities in the same unit unit_to_entities = {} for unit_id, entity_id in unit_entity_pairs: if unit_id not in unit_to_entities: unit_to_entities[unit_id] = set() unit_to_entities[unit_id].add(entity_id) # Update co-occurrences for all pairs in each unit cooccurrence_pairs = set() # Use set to avoid duplicates for unit_id, entity_ids in unit_to_entities.items(): entity_list = list(entity_ids) # Convert set to list for iteration # For each pair of entities in this unit, create co-occurrence for i, entity_id_1 in enumerate(entity_list): for entity_id_2 in entity_list[i+1:]: # Skip if same entity (shouldn't happen with set, but be safe) if entity_id_1 == entity_id_2: continue # Ensure consistent ordering (entity_id_1 < entity_id_2) if entity_id_1 > entity_id_2: entity_id_1, entity_id_2 = entity_id_2, entity_id_1 cooccurrence_pairs.add((entity_id_1, entity_id_2)) # Batch update co-occurrences if cooccurrence_pairs: now = datetime.now(timezone.utc) await conn.executemany( """ INSERT INTO entity_cooccurrences (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred) VALUES ($1, $2, $3, $4) ON CONFLICT (entity_id_1, entity_id_2) DO UPDATE SET cooccurrence_count = entity_cooccurrences.cooccurrence_count + 1, last_cooccurred = EXCLUDED.last_cooccurred """, [(e1, e2, 1, now) for e1, e2 in cooccurrence_pairs] ) async def get_units_by_entity(self, entity_id: str, limit: int = 100) -> List[str]: """ Get all units that mention an entity. Args: entity_id: Entity ID limit: Max results Returns: List of unit IDs """ async with self.pool.acquire() as conn: rows = await conn.fetch( """ SELECT unit_id FROM unit_entities WHERE entity_id = $1 ORDER BY unit_id LIMIT $2 """, entity_id, limit ) return [row['unit_id'] for row in rows] async def get_entity_by_text( self, agent_id: str, entity_text: str, entity_type: Optional[str] = None ) -> Optional[str]: """ Find an entity by text (for query resolution). Args: agent_id: Agent ID entity_text: Entity text to search for entity_type: Optional entity type filter Returns: Entity ID if found, None otherwise """ async with self.pool.acquire() as conn: if entity_type: row = await conn.fetchrow( """ SELECT id FROM entities WHERE agent_id = $1 AND entity_type = $2 AND canonical_name ILIKE $3 ORDER BY mention_count DESC LIMIT 1 """, agent_id, entity_type, entity_text ) else: row = await conn.fetchrow( """ SELECT id FROM entities WHERE agent_id = $1 AND canonical_name ILIKE $2 ORDER BY mention_count DESC LIMIT 1 """, agent_id, entity_text ) return row['id'] if row else None