# conversation state tracking for entity and schema reference persistence import logging import re import time from typing import Optional, Dict, Any, List, Set from dataclasses import dataclass, field from collections import OrderedDict CONVERSATION_STATE_MAX_ENTITIES = 50 logger = logging.getLogger(__name__) @dataclass class EntityReference: # tracked entity from conversation name: str entity_type: str value: Any = None source_turn: int = 0 confidence: float = 1.0 timestamp: float = 0.0 def to_dict(self) -> Dict[str, Any]: return { "name": self.name, "entity_type": self.entity_type, "value": self.value, "source_turn": self.source_turn, "confidence": self.confidence, } @dataclass class QueryRecord: # record of a successfully executed query query_text: str sql: str tables: List[str] columns: List[str] filters: Dict[str, Any] row_count: int turn_number: int timestamp: float = 0.0 def to_dict(self) -> Dict[str, Any]: return { "query_text": self.query_text, "sql": self.sql, "tables": self.tables, "columns": self.columns, "row_count": self.row_count, "turn_number": self.turn_number, } class ConversationStateManager: # tracks entities, tables, and schema references across conversation turns MAX_QUERY_RECORDS = 10 MAX_ENTITIES = CONVERSATION_STATE_MAX_ENTITIES def __init__(self): self._entities: OrderedDict[str, EntityReference] = OrderedDict() self._query_records: List[QueryRecord] = [] self._referenced_tables: OrderedDict[str, int] = OrderedDict() self._turn_count: int = 0 self._active_filters: Dict[str, Any] = {} @property def turn_count(self) -> int: return self._turn_count @property def referenced_tables(self) -> List[str]: # returns tables in order of most recent reference return list(reversed(self._referenced_tables.keys())) @property def active_filters(self) -> Dict[str, Any]: return self._active_filters.copy() @property def last_query(self) -> Optional[QueryRecord]: return self._query_records[-1] if self._query_records else None @property def last_tables(self) -> List[str]: # returns tables from the most recent query if self._query_records: return self._query_records[-1].tables return [] @property def last_sql(self) -> Optional[str]: if self._query_records: return self._query_records[-1].sql return None def advance_turn(self) -> None: # increments conversation turn counter self._turn_count += 1 def record_query( self, query_text: str, sql: str, tables: List[str], columns: Optional[List[str]] = None, filters: Optional[Dict[str, Any]] = None, row_count: int = 0, ) -> None: # records a successful query execution self.advance_turn() record = QueryRecord( query_text=query_text, sql=sql, tables=tables, columns=columns or [], filters=filters or {}, row_count=row_count, turn_number=self._turn_count, timestamp=time.time(), ) self._query_records.append(record) # prune old records if len(self._query_records) > self.MAX_QUERY_RECORDS: self._query_records = self._query_records[-self.MAX_QUERY_RECORDS:] # update referenced tables for table in tables: self._referenced_tables[table] = self._turn_count self._referenced_tables.move_to_end(table) # update active filters if filters: self._active_filters.update(filters) # extract entities from query self._extract_entities(query_text, tables) def add_entity(self, name: str, entity_type: str, value: Any = None, confidence: float = 1.0) -> None: # adds or updates an entity reference self._entities[name] = EntityReference( name=name, entity_type=entity_type, value=value, source_turn=self._turn_count, confidence=confidence, timestamp=time.time(), ) self._entities.move_to_end(name) # prune if too many while len(self._entities) > self.MAX_ENTITIES: self._entities.popitem(last=False) def get_entity(self, name: str) -> Optional[EntityReference]: return self._entities.get(name) def get_entities_by_type(self, entity_type: str) -> List[EntityReference]: return [e for e in self._entities.values() if e.entity_type == entity_type] def get_recent_queries(self, count: int = 3) -> List[QueryRecord]: return self._query_records[-count:] def get_context_summary(self) -> str: # builds context summary string for LLM prompt injection parts = [] if self._query_records: parts.append("ÖNCEKİ SORGULAR:") for record in self._query_records[-3:]: parts.append(f" Soru: {record.query_text[:100]}") parts.append(f" SQL: {record.sql[:200]}") parts.append(f" Tablolar: {', '.join(record.tables)}") parts.append(f" Sonuç: {record.row_count} kayıt") parts.append("") if self._referenced_tables: parts.append(f"KULLANILAN TABLOLAR: {', '.join(self.referenced_tables[:10])}") if self._active_filters: filters_str = ", ".join(f"{k}={v}" for k, v in list(self._active_filters.items())[:5]) parts.append(f"AKTIF FILTRELER: {filters_str}") return "\n".join(parts) def resolve_table_reference(self, text: str) -> Optional[str]: # resolves pronoun-like references to tables (e.g. "bunlar" -> last table) reference_patterns = { "bunlar": -1, "bu kayitlar": -1, "bu sonuclar": -1, "onceki": -1, "oncekiler": -1, "yukardaki": -1, "onlar": -1, "bu dosyalar": -1, "bu veriler": -1, } text_lower = text.lower() for pattern, offset in reference_patterns.items(): if pattern in text_lower: if self._query_records: idx = max(0, len(self._query_records) + offset) if idx < len(self._query_records): tables = self._query_records[idx].tables return tables[0] if tables else None return None def get_follow_up_context(self) -> Dict[str, Any]: # returns context data useful for resolving follow-up queries return { "last_tables": self.last_tables, "last_sql": self.last_sql, "referenced_tables": self.referenced_tables[:5], "active_filters": self._active_filters, "turn_count": self._turn_count, "recent_queries": [r.to_dict() for r in self.get_recent_queries(3)], } def _extract_entities(self, query: str, tables: List[str]) -> None: # extracts table and numeric entities from query text for table in tables: self.add_entity(table, "table", table) # extract numeric values that might be IDs or filter values numbers = re.findall(r'\b(\d+)\b', query) for num in numbers[:5]: self.add_entity(f"value_{num}", "numeric", int(num), confidence=0.6) def clear(self) -> None: # resets all conversation state self._entities.clear() self._query_records.clear() self._referenced_tables.clear() self._turn_count = 0 self._active_filters.clear() def to_dict(self) -> Dict[str, Any]: return { "turn_count": self._turn_count, "entities": {k: v.to_dict() for k, v in self._entities.items()}, "query_records": [r.to_dict() for r in self._query_records], "referenced_tables": list(self._referenced_tables.keys()), "active_filters": self._active_filters, } # module-level singleton _state_instance: Optional[ConversationStateManager] = None def get_conversation_state() -> ConversationStateManager: # get or create singleton conversation state global _state_instance if _state_instance is None: _state_instance = ConversationStateManager() return _state_instance def reset_conversation_state() -> None: # reset singleton for testing global _state_instance if _state_instance is not None: _state_instance.clear() _state_instance = None