diff --git a/src/deriver/consumer.py b/src/deriver/consumer.py index 399904d7..4420aaa5 100644 --- a/src/deriver/consumer.py +++ b/src/deriver/consumer.py @@ -5,7 +5,6 @@ import sentry_sdk from langfuse.decorators import langfuse_context from pydantic import ValidationError from rich.console import Console -from sqlalchemy.ext.asyncio import AsyncSession from src.config import settings from src.dependencies import tracked_db @@ -45,36 +44,41 @@ async def process_item(task_type: str, payload: dict[str, Any]) -> None: raise ValueError(f"Invalid payload structure: {str(e)}") from e await process_webhook(validated) logger.debug("Finished processing webhook %s", validated.event_type) + elif task_type == "summary": + if settings.LANGFUSE_PUBLIC_KEY: + langfuse_context.update_current_trace( # type: ignore + metadata={ + "critical_analysis_model": settings.DERIVER.MODEL, + } + ) + try: + validated = SummaryPayload(**payload) + except ValidationError as e: + logger.error( + "Invalid summary payload received: %s. Payload: %s", str(e), payload + ) + raise ValueError(f"Invalid payload structure: {str(e)}") from e + await process_summary_task(validated) + elif task_type == "representation": + if settings.LANGFUSE_PUBLIC_KEY: + langfuse_context.update_current_trace( + metadata={ + "critical_analysis_model": settings.DERIVER.MODEL, + } + ) - if settings.LANGFUSE_PUBLIC_KEY: - langfuse_context.update_current_trace( - metadata={ - "critical_analysis_model": settings.DERIVER.MODEL, - } - ) - - # Open a DB session only for the duration of the processing call - async with tracked_db("deriver") as db: - if task_type == "summary": - try: - validated = SummaryPayload(**payload) - except ValidationError as e: - logger.error( - "Invalid summary payload received: %s. Payload: %s", str(e), payload - ) - raise ValueError(f"Invalid payload structure: {str(e)}") from e - await process_summary_task(db, validated) - elif task_type == "representation": - try: - validated = RepresentationPayload(**payload) - except ValidationError as e: - logger.error( - "Invalid representation payload received: %s. Payload: %s", - str(e), - payload, - ) - raise ValueError(f"Invalid payload structure: {str(e)}") from e - await deriver.process_representation_task(db, validated) + try: + validated = RepresentationPayload(**payload) + except ValidationError as e: + logger.error( + "Invalid representation payload received: %s. Payload: %s", + str(e), + payload, + ) + raise ValueError(f"Invalid payload structure: {str(e)}") from e + await deriver.process_representation_task(validated) + else: + raise ValueError(f"Invalid task type: {task_type}") @sentry_sdk.trace @@ -87,14 +91,12 @@ async def process_webhook( @sentry_sdk.trace async def process_summary_task( - db: AsyncSession, payload: SummaryPayload, ) -> None: """ Process a summary task by generating summaries if needed. """ await summarizer.summarize_if_needed( - db, payload.workspace_name, payload.session_name, payload.message_id, diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index dd43a0ac..037ac864 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -6,11 +6,11 @@ from typing import Any import sentry_sdk from langfuse.decorators import langfuse_context -from sqlalchemy.ext.asyncio import AsyncSession from src import crud, exceptions from src.config import settings from src.crud.representation import GLOBAL_REPRESENTATION_COLLECTION_NAME +from src.dependencies import tracked_db from src.utils import summarizer from src.utils.clients import honcho_llm_call from src.utils.embedding_store import EmbeddingStore @@ -105,7 +105,6 @@ async def peer_card_call( @conditional_observe @sentry_sdk.trace async def process_representation_task( - db: AsyncSession, payload: RepresentationPayload, ) -> None: """ @@ -117,14 +116,15 @@ async def process_representation_task( logger.debug("Starting insight extraction for user message: %s", payload.message_id) # Use get_session_context_formatted with configurable token limit - formatted_history = await summarizer.get_session_context_formatted( - db, - payload.workspace_name, - payload.session_name, - token_limit=settings.DERIVER.CONTEXT_TOKEN_LIMIT, - cutoff=payload.message_id, - include_summary=True, - ) + async with tracked_db("deriver.get_session_context") as db: + formatted_history = await summarizer.get_session_context_formatted( + db, + payload.workspace_name, + payload.session_name, + token_limit=settings.DERIVER.CONTEXT_TOKEN_LIMIT, + cutoff=payload.message_id, + include_summary=True, + ) # instantiate embedding store from collection # if the sender is also the target, we're handling a global representation task. @@ -139,33 +139,36 @@ async def process_representation_task( ) # get_or_create_collection already handles IntegrityError with rollback and a retry - collection = await crud.get_or_create_collection( - db, - payload.workspace_name, - collection_name, - payload.sender_name, - ) + async with tracked_db("deriver.get_or_create_collection") as db: + collection = await crud.get_or_create_collection( + db, + payload.workspace_name, + collection_name, + payload.sender_name, + ) + collection_name_loaded = collection.name # Use the embedding store directly embedding_store = EmbeddingStore( workspace_name=payload.workspace_name, peer_name=payload.sender_name, - collection_name=collection.name, + collection_name=collection_name_loaded, ) # Create reasoner instance reasoner = CertaintyReasoner(embedding_store=embedding_store, ctx=payload) # Check for existing working representation first, fall back to global search - working_rep_data: ( - dict[str, Any] | str | None - ) = await crud.get_working_representation_data( - db, - payload.workspace_name, - payload.target_name, - payload.sender_name, - payload.session_name, - ) + async with tracked_db("deriver.get_working_representation_data") as db: + working_rep_data: ( + dict[str, Any] | str | None + ) = await crud.get_working_representation_data( + db, + payload.workspace_name, + payload.target_name, + payload.sender_name, + payload.session_name, + ) # Time context preparation context_prep_start = time.perf_counter() @@ -222,9 +225,10 @@ async def process_representation_task( # We currently only use Peer Cards in Honcho-level representation derivation. if payload.sender_name == payload.target_name: - sender_peer_card: list[str] | None = await crud.get_peer_card( - db, payload.workspace_name, payload.sender_name - ) + async with tracked_db("deriver.get_peer_card") as db: + sender_peer_card: list[str] | None = await crud.get_peer_card( + db, payload.workspace_name, payload.sender_name + ) if sender_peer_card is None: logger.warning("No peer card found for %s", payload.sender_name) else: @@ -235,7 +239,6 @@ async def process_representation_task( # Run single-pass reasoning final_observations = await reasoner.reason( - db, working_representation, formatted_history, sender_peer_card, @@ -250,7 +253,7 @@ async def process_representation_task( log_observations_tree(final_obs_dict) # Always save working representation to peer for dialectic access - await save_working_representation_to_peer(db, payload, final_observations) + await save_working_representation_to_peer(payload, final_observations) # Calculate and log overall timing overall_duration = (time.perf_counter() - overall_start) * 1000 @@ -404,7 +407,6 @@ class CertaintyReasoner: @sentry_sdk.trace async def reason( self, - db: AsyncSession, working_representation: ReasoningResponseWithThinking, history: str, speaker_peer_card: list[str] | None, @@ -460,7 +462,7 @@ class CertaintyReasoner: for observation in level ] if new_observations: - await self._update_peer_card(db, speaker_peer_card, new_observations) + await self._update_peer_card(speaker_peer_card, new_observations) update_peer_card_duration = ( time.perf_counter() - update_peer_card_start ) * 1000 @@ -539,7 +541,6 @@ class CertaintyReasoner: @sentry_sdk.trace async def _update_peer_card( self, - db: AsyncSession, old_peer_card: list[str] | None, new_observations: list[str], ) -> None: @@ -554,9 +555,10 @@ class CertaintyReasoner: logger.info("No changes to peer card") return logger.info("New peer card: %s", new_peer_card) - await crud.set_peer_card( - db, self.ctx.workspace_name, self.ctx.sender_name, new_peer_card - ) + async with tracked_db("deriver.update_peer_card") as db: + await crud.set_peer_card( + db, self.ctx.workspace_name, self.ctx.sender_name, new_peer_card + ) except Exception as e: if settings.SENTRY.ENABLED: sentry_sdk.capture_exception(e) @@ -592,7 +594,6 @@ def observation_context_to_reasoning_response( @sentry_sdk.trace async def save_working_representation_to_peer( - db: AsyncSession, payload: RepresentationPayload, final_observations: ReasoningResponseWithThinking, ) -> None: @@ -617,11 +618,12 @@ async def save_working_representation_to_peer( "created_at": utc_now_iso(), } - await crud.set_working_representation( - db, - working_rep_data, - payload.workspace_name, - payload.target_name, - payload.sender_name, - payload.session_name, - ) + async with tracked_db("deriver.save_working_representation") as db: + await crud.set_working_representation( + db, + working_rep_data, + payload.workspace_name, + payload.target_name, + payload.sender_name, + payload.session_name, + ) diff --git a/src/utils/summarizer.py b/src/utils/summarizer.py index 8c8394cd..c2337337 100644 --- a/src/utils/summarizer.py +++ b/src/utils/summarizer.py @@ -174,7 +174,6 @@ Produce as thorough a summary as possible in {output_words} words or less. async def summarize_if_needed( - db: AsyncSession, workspace_name: str, session_name: str, message_id: int, @@ -187,7 +186,6 @@ async def summarize_if_needed( without assuming any relationship between their thresholds. Args: - db: Database session workspace_name: The workspace name session_name: The session name message_id: The message ID @@ -239,35 +237,36 @@ async def summarize_if_needed( return_exceptions=True, ) else: - # If only one summary needs to be created, run them individually - if should_create_long: - await _create_and_save_summary( - db, - workspace_name, - session_name, - message_id, - SummaryType.LONG, - ) - logger.info( - "Saved long summary for session %s covering up to message %s (%s in session)", - session_name, - message_id, - message_seq_in_session, - ) - elif should_create_short: - await _create_and_save_summary( - db, - workspace_name, - session_name, - message_id, - SummaryType.SHORT, - ) - logger.info( - "Saved short summary for session %s covering up to message %s (%s in session)", - session_name, - message_id, - message_seq_in_session, - ) + async with tracked_db("create_summary") as db: + # If only one summary needs to be created, run them individually + if should_create_long: + await _create_and_save_summary( + db, + workspace_name, + session_name, + message_id, + SummaryType.LONG, + ) + logger.info( + "Saved long summary for session %s covering up to message %s (%s in session)", + session_name, + message_id, + message_seq_in_session, + ) + elif should_create_short: + await _create_and_save_summary( + db, + workspace_name, + session_name, + message_id, + SummaryType.SHORT, + ) + logger.info( + "Saved short summary for session %s covering up to message %s (%s in session)", + session_name, + message_id, + message_seq_in_session, + ) async def _create_and_save_summary(