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:
parent
9c8e74bc04
commit
ccaffbba67
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Reference in New Issue