diff --git a/src/agent.py b/src/agent.py index 21ed2698..bb53aa90 100644 --- a/src/agent.py +++ b/src/agent.py @@ -43,23 +43,6 @@ QUERY_GENERATION_SYSTEM = """Given this query about a user, generate 3 focused s load_dotenv() -class AsyncSet: - def __init__(self): - self._set: set[str] = set() - self._lock = asyncio.Lock() - - async def add(self, item: str): - async with self._lock: - self._set.add(item) - - async def update(self, items: Iterable[str]): - async with self._lock: - self._set.update(items) - - def get_set(self) -> set[str]: - return self._set.copy() - - class Dialectic: def __init__(self, agent_input: str, user_representation: str, chat_history: str): self.agent_input = agent_input @@ -152,7 +135,6 @@ async def chat( user_id: str, session_id: str, queries: str | list[str], - db: AsyncSession, stream: bool = False, ) -> schemas.DialecticResponse | MessageStreamManager: """ @@ -178,69 +160,65 @@ async def chat( # Setup phase - create resources we'll need for all operations - # 1. Create embedding store - collection = await crud.get_or_create_user_protected_collection(db, app_id, user_id) + # 1. Fetch latest user message & chat history + async with tracked_db("chat.load_history") as db_history: + stmt = ( + select(models.Message) + .where(models.Message.app_id == app_id) + .where(models.Message.user_id == user_id) + .where(models.Message.session_id == session_id) + .where(models.Message.is_user) + .order_by(models.Message.id.desc()) + .limit(1) + ) + latest_messages = await db_history.execute(stmt) + latest_message = latest_messages.scalar_one_or_none() + latest_message_id = latest_message.public_id if latest_message else None + logger.debug(f"Latest user message ID: {latest_message_id}") - embedding_store = CollectionEmbeddingStore( - db=db, - app_id=app_id, - user_id=user_id, - collection_id=collection.public_id, # type: ignore - ) - logger.debug( - f"Created embedding store with collection_id: {collection.public_id if collection else None}" - ) + chat_history, _, _ = await history.get_summarized_history( + db_history, session_id, summary_type=history.SummaryType.SHORT + ) + if not chat_history: + logger.warning(f"No chat history found for session {session_id}") + chat_history = f"someone asked this about the user's message: {final_query}" + logger.debug(f"IDs: {app_id}, {user_id}, {session_id}") + message_count = len(chat_history.split("\n")) + logger.debug(f"Retrieved chat history: {message_count} messages") - stmt = ( - select(models.Message) - .where(models.Message.app_id == app_id) - .where(models.Message.user_id == user_id) - .where(models.Message.session_id == session_id) - .where(models.Message.is_user) - .order_by(models.Message.id.desc()) - .limit(1) - ) + # Run short-term inference and long-term facts in parallel + async def fetch_long_term(): + async with tracked_db("chat.get_collection") as db_embed: + collection = await crud.get_or_create_user_protected_collection( + db_embed, app_id, user_id + ) + collection_id = ( + collection.public_id + ) # Extract the ID while session is active + facts = await get_long_term_facts(final_query, app_id, user_id, collection_id) + return facts - latest_messages = await db.execute(stmt) - latest_message = latest_messages.scalar_one_or_none() - latest_message_id = latest_message.public_id if latest_message else None - logger.debug(f"Latest user message ID: {latest_message_id}") + long_term_task = asyncio.create_task(fetch_long_term()) + short_term_task = asyncio.create_task(run_tom_inference(chat_history, session_id)) - # Get chat history for the session - chat_history, _, _ = await history.get_summarized_history( - db, session_id, summary_type=history.SummaryType.SHORT - ) - if not chat_history: - logger.warning(f"No chat history found for session {session_id}") - chat_history = f"someone asked this about the user's message: {final_query}" - logger.debug(f"IDs: {app_id}, {user_id}, {session_id}") - message_count = len(chat_history.split("\n")) - logger.debug(f"Retrieved chat history: {message_count} messages") - - # Run both long-term and short-term context retrieval concurrently - logger.debug("Starting parallel tasks for context retrieval") - long_term_task = get_long_term_facts(final_query, embedding_store) - short_term_task = run_tom_inference(chat_history, session_id) - - # Wait for both tasks to complete facts, tom_inference = await asyncio.gather(long_term_task, short_term_task) logger.debug(f"Retrieved {len(facts)} facts from long-term memory") logger.debug(f"TOM inference completed with {len(tom_inference)} characters") # Generate a fresh user representation logger.debug("Generating user representation") - user_representation = await generate_user_representation( - app_id=app_id, - user_id=user_id, - session_id=session_id, - chat_history=chat_history, - tom_inference=tom_inference, - facts=facts, - embedding_store=embedding_store, - db=db, - message_id=latest_message_id, - with_inference=False, - ) + async with tracked_db("chat.generate_user_representation") as db_rep: + user_representation = await generate_user_representation( + app_id=app_id, + user_id=user_id, + session_id=session_id, + chat_history=chat_history, + tom_inference=tom_inference, + facts=facts, + db=db_rep, + message_id=latest_message_id, + with_inference=False, + ) logger.debug( f"User representation generated: {len(user_representation)} characters" ) @@ -282,7 +260,7 @@ async def chat( async def get_long_term_facts( - query: str, embedding_store: CollectionEmbeddingStore + query: str, app_id: str, user_id: str, collection_id: str ) -> list[str]: """ Generate queries based on the dialectic query and retrieve relevant facts. @@ -306,7 +284,12 @@ async def get_long_term_facts( async def execute_query(i: int, search_query: str) -> list[str]: logger.debug(f"Starting query {i + 1}/{len(search_queries)}: {search_query}") query_start = asyncio.get_event_loop().time() - facts = await embedding_store.get_relevant_facts( + query_embedding_store = CollectionEmbeddingStore( + app_id=app_id, + user_id=user_id, + collection_id=collection_id, + ) + facts = await query_embedding_store.get_relevant_facts( search_query, top_k=10, max_distance=0.85 ) query_time = asyncio.get_event_loop().time() - query_start @@ -430,7 +413,6 @@ async def generate_user_representation( chat_history: str, tom_inference: str, facts: list[str], - embedding_store: CollectionEmbeddingStore, db: AsyncSession, message_id: Optional[str] = None, with_inference: bool = False, @@ -477,7 +459,6 @@ async def generate_user_representation( chat_history=chat_history, session_id=session_id, facts=facts, - embedding_store=embedding_store, user_representation=latest_representation, tom_inference=tom_inference, ) @@ -505,36 +486,32 @@ RELEVANT LONG-TERM FACTS ABOUT THE USER: logger.debug(f"Saving representation to message_id: {message_id}") save_start = asyncio.get_event_loop().time() try: - async with tracked_db("agent.generate_user_representation") as save_db: - try: - # First check if message exists - message_check_stmt = select(models.Message).where( - models.Message.public_id == message_id - ) - message_check = await save_db.execute(message_check_stmt) - message_exists = message_check.scalar_one_or_none() is not None + # First check if message exists + message_check_stmt = select(models.Message).where( + models.Message.public_id == message_id + ) + message_check = await db.execute(message_check_stmt) + message_exists = message_check.scalar_one_or_none() is not None - if not message_exists: - message_id = None - else: - metamessage = models.Metamessage( - app_id=app_id, - user_id=user_id, - session_id=session_id, - message_id=message_id if message_id else None, - label=USER_REPRESENTATION_METAMESSAGE_TYPE, - content=representation, - h_metadata={}, - ) - save_db.add(metamessage) - await save_db.commit() - save_time = asyncio.get_event_loop().time() - save_start - logger.debug(f"Representation saved in {save_time:.2f}s") - except Exception as inner_e: - logger.error(f"Error during save DB operation: {str(inner_e)}") - await save_db.rollback() + if not message_exists: + message_id = None + else: + metamessage = models.Metamessage( + app_id=app_id, + user_id=user_id, + session_id=session_id, + message_id=message_id if message_id else None, + label=USER_REPRESENTATION_METAMESSAGE_TYPE, + content=representation, + h_metadata={}, + ) + db.add(metamessage) + await db.commit() + save_time = asyncio.get_event_loop().time() - save_start + logger.debug(f"Representation saved in {save_time:.2f}s") except Exception as e: - logger.error(f"Error creating DB session: {str(e)}") + logger.error(f"Error during save DB operation: {str(e)}") + await db.rollback() total_time = asyncio.get_event_loop().time() - rep_start_time logger.debug(f"Total representation generation completed in {total_time:.2f}s") diff --git a/src/deriver/consumer.py b/src/deriver/consumer.py index 50dce270..cc6525d0 100644 --- a/src/deriver/consumer.py +++ b/src/deriver/consumer.py @@ -119,7 +119,6 @@ async def process_user_message( db=db, app_id=app_id, user_id=user_id ) embedding_store = CollectionEmbeddingStore( - db=db, app_id=app_id, user_id=user_id, collection_id=collection.public_id, # type: ignore diff --git a/src/deriver/tom/embeddings.py b/src/deriver/tom/embeddings.py index 031c28c1..059ea791 100644 --- a/src/deriver/tom/embeddings.py +++ b/src/deriver/tom/embeddings.py @@ -3,13 +3,13 @@ import logging from sqlalchemy.ext.asyncio import AsyncSession from ... import crud, schemas +from ...dependencies import tracked_db logger = logging.getLogger(__name__) class CollectionEmbeddingStore: - def __init__(self, db: AsyncSession, app_id: str, user_id: str, collection_id: str): - self.db = db + def __init__(self, app_id: str, user_id: str, collection_id: str): self.app_id = app_id self.user_id = user_id self.collection_id = collection_id @@ -28,24 +28,25 @@ class CollectionEmbeddingStore: replace_duplicates: If True, replace old duplicates with new facts. If False, discard new duplicates similarity_threshold: Facts with similarity above this threshold are considered duplicates """ - for fact in facts: - # Create document with duplicate checking - try: - metadata = {} - if message_id: - metadata["message_id"] = message_id - await crud.create_document( - self.db, - document=schemas.DocumentCreate(content=fact, metadata=metadata), - app_id=self.app_id, - user_id=self.user_id, - collection_id=self.collection_id, - duplicate_threshold=1 - - similarity_threshold, # Convert similarity to distance - ) - except Exception as e: - logger.error(f"Error creating document: {e}") - continue + async with tracked_db("embedding_store.save_facts") as db: + for fact in facts: + # Create document with duplicate checking + try: + metadata = {} + if message_id: + metadata["message_id"] = message_id + await crud.create_document( + db, + document=schemas.DocumentCreate(content=fact, metadata=metadata), + app_id=self.app_id, + user_id=self.user_id, + collection_id=self.collection_id, + duplicate_threshold=1 + - similarity_threshold, # Convert similarity to distance + ) + except Exception as e: + logger.error(f"Error creating document: {e}") + continue async def get_relevant_facts( self, query: str, top_k: int = 5, max_distance: float = 0.3 @@ -60,17 +61,18 @@ class CollectionEmbeddingStore: Returns: List of facts sorted by relevance """ - documents = await crud.query_documents( - self.db, - app_id=self.app_id, - user_id=self.user_id, - collection_id=self.collection_id, - query=query, - max_distance=max_distance, - top_k=top_k, - ) + async with tracked_db("embedding_store.get_relevant_facts") as db: + documents = await crud.query_documents( + db, + app_id=self.app_id, + user_id=self.user_id, + collection_id=self.collection_id, + query=query, + max_distance=max_distance, + top_k=top_k, + ) - return [doc.content for doc in documents] + return [doc.content for doc in documents] async def remove_duplicates( self, facts: list[str], similarity_threshold: float = 0.85 @@ -86,29 +88,30 @@ class CollectionEmbeddingStore: """ unique_facts = [] - for fact in facts: - try: - # Check for duplicates using the crud function - duplicates = await crud.get_duplicate_documents( - self.db, - app_id=self.app_id, - user_id=self.user_id, - collection_id=self.collection_id, - content=fact, - similarity_threshold=similarity_threshold, - ) - - if not duplicates: - # No duplicates found, add to unique facts - unique_facts.append(fact) - else: - # Log duplicate found - logger.debug( - f"Duplicate found: {duplicates[0].content}. Ignoring fact: {fact}" + async with tracked_db("embedding_store.remove_duplicates") as db: + for fact in facts: + try: + # Check for duplicates using the crud function + duplicates = await crud.get_duplicate_documents( + db, + app_id=self.app_id, + user_id=self.user_id, + collection_id=self.collection_id, + content=fact, + similarity_threshold=similarity_threshold, ) - except Exception as e: - logger.error(f"Error checking for duplicates: {e}") - # If there's an error, still include the fact to avoid losing information - unique_facts.append(fact) + + if not duplicates: + # No duplicates found, add to unique facts + unique_facts.append(fact) + else: + # Log duplicate found + logger.debug( + f"Duplicate found: {duplicates[0].content}. Ignoring fact: {fact}" + ) + except Exception as e: + logger.error(f"Error checking for duplicates: {e}") + # If there's an error, still include the fact to avoid losing information + unique_facts.append(fact) return unique_facts diff --git a/src/deriver/tom/long_term.py b/src/deriver/tom/long_term.py index c40fcbcf..a9a7f2c5 100644 --- a/src/deriver/tom/long_term.py +++ b/src/deriver/tom/long_term.py @@ -9,8 +9,6 @@ from sentry_sdk.ai.monitoring import ai_track from src.utils import parse_xml_content from src.utils.model_client import ModelClient, ModelProvider -from .embeddings import CollectionEmbeddingStore - # Configure logging logger = logging.getLogger(__name__) @@ -29,7 +27,6 @@ MAX_FACT_DISTANCE = 0.85 async def get_user_representation_long_term( chat_history: str, session_id: str, - embedding_store: CollectionEmbeddingStore, user_representation: str = "None", tom_inference: str = "None", facts: Optional[list[str]] = None, diff --git a/src/routers/sessions.py b/src/routers/sessions.py index 51819332..b7108251 100644 --- a/src/routers/sessions.py +++ b/src/routers/sessions.py @@ -216,7 +216,6 @@ async def chat( options: schemas.DialecticOptions = Body( ..., description="Dialectic Endpoint Parameters" ), - db=db, ): """Chat with the Dialectic API""" @@ -226,7 +225,6 @@ async def chat( user_id=user_id, session_id=session_id, queries=options.queries, - db=db, ) else: @@ -238,7 +236,6 @@ async def chat( session_id=session_id, queries=options.queries, stream=True, - db=db, ) if type(stream) is AsyncMessageStreamManager: async with stream as stream_manager: