refactor: scope tracked_db usage within process_representation_task (#194)

* refactor: scope tracked_db usage within process_representation_task

* fix: clean up method

* fix: load collection name
This commit is contained in:
Rajat Ahuja 2025-08-14 13:15:25 -04:00 committed by GitHub
parent 9c8e74bc04
commit ccaffbba67
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 112 additions and 109 deletions

View File

@ -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,

View File

@ -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,
)

View File

@ -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(