honcho/src/crud/representation.py

624 lines
24 KiB
Python

from __future__ import annotations
import datetime
import logging
import time
from contextlib import suppress
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from src import crud, exceptions, models, schemas
from src.config import settings
from src.dependencies import tracked_db
from src.dreamer.dream_scheduler import check_and_schedule_dream
from src.embedding_client import EmbeddingTokenLimitError, embedding_client
from src.schemas import ResolvedConfiguration
from src.telemetry.events import EmbeddingCallPurpose
from src.telemetry.logging import accumulate_metric
from src.utils.formatting import format_datetime_utc
from src.utils.representation import (
ALLOWLIST_SAFE_LEVELS,
DeductiveObservation,
ExplicitObservation,
Representation,
allowlist_safe_levels,
)
from src.utils.types import embedding_call_purpose
logger = logging.getLogger(__name__)
def _observation_text(obs: ExplicitObservation | DeductiveObservation) -> str:
"""Return the canonical text payload for an explicit or deductive observation."""
return obs.conclusion if isinstance(obs, DeductiveObservation) else obs.content
def _normalized_observation(
obs: ExplicitObservation | DeductiveObservation,
) -> ExplicitObservation | DeductiveObservation:
"""Return an observation with its persisted/embed text normalized."""
text = _observation_text(obs).strip()
if isinstance(obs, DeductiveObservation):
return obs.model_copy(update={"conclusion": text})
return obs.model_copy(update={"content": text})
class RepresentationManager:
"""Unified manager for representation and document queries."""
def __init__(
self,
workspace_name: str,
*,
observer: str,
observed: str,
) -> None:
self.workspace_name: str = workspace_name
self.observer: str = observer
self.observed: str = observed
async def save_representation(
self,
representation: Representation,
message_ids: list[int],
session_name: str,
message_created_at: datetime.datetime,
message_level_configuration: ResolvedConfiguration,
) -> crud.CreateDocumentsResult:
"""
Save Representation objects to the collection as a set of documents.
Args:
representation: Representation object
message_ids: Message ID range to link with observations
session_name: Session name to link with existing summary context
message_created_at: Timestamp when the message was created
Returns:
The result of document creation, including saved documents and
deduplication counts.
"""
empty_result = crud.CreateDocumentsResult()
if not representation.deductive and not representation.explicit:
logger.debug("No observations to save")
return empty_result
all_observations = [
_normalized_observation(obs)
for obs in representation.deductive + representation.explicit
if _observation_text(obs).strip()
]
if not all_observations:
logger.debug("No non-empty observations to save")
return empty_result
# Batch embed all observations
batch_embed_start = time.perf_counter()
observation_texts = [_observation_text(obs) for obs in all_observations]
try:
with embedding_call_purpose(
EmbeddingCallPurpose.CREATE_OBSERVATIONS.value,
workspace_name=self.workspace_name,
parent_category="representation",
):
embeddings = await embedding_client.simple_batch_embed(
observation_texts, on_oversize="truncate"
)
except EmbeddingTokenLimitError as e:
raise exceptions.ValidationException(
"Observation content exceeds maximum token limit of "
+ f"{settings.EMBEDDING.MAX_INPUT_TOKENS}."
) from e
batch_embed_duration = (time.perf_counter() - batch_embed_start) * 1000
accumulate_metric(
f"deriver_{message_ids[-1]}_{self.observer}",
"embed_new_observations",
batch_embed_duration,
"ms",
)
# Batch create document objects
create_document_start = time.perf_counter()
async with tracked_db("representation_manager.save_representation") as db:
new_documents_result = await self._save_representation_internal(
db,
all_observations,
embeddings,
message_ids,
session_name,
message_created_at,
message_level_configuration,
)
create_document_duration = (time.perf_counter() - create_document_start) * 1000
accumulate_metric(
f"deriver_{message_ids[-1]}_{self.observer}",
"save_new_observations",
create_document_duration,
"ms",
)
return new_documents_result
async def _save_representation_internal(
self,
db: AsyncSession,
all_observations: list[ExplicitObservation | DeductiveObservation],
embeddings: list[list[float]],
message_ids: list[int],
session_name: str,
message_created_at: datetime.datetime,
message_level_configuration: ResolvedConfiguration,
) -> crud.CreateDocumentsResult:
# get_or_create_collection already handles IntegrityError with rollback and a retry
collection = await crud.get_or_create_collection(
db,
self.workspace_name,
observer=self.observer,
observed=self.observed,
)
# Prepare all documents for bulk creation
documents_to_create: list[schemas.DocumentCreate] = []
for obs, embedding in zip(all_observations, embeddings, strict=True):
# NOTE: will add additional levels of reasoning in the future
if isinstance(obs, DeductiveObservation):
obs_level = "deductive"
obs_content = obs.conclusion
obs_premises = obs.premises
else:
obs_level = "explicit"
obs_content = obs.content
obs_premises = None
metadata: schemas.DocumentMetadata = schemas.DocumentMetadata(
message_ids=message_ids,
premises=obs_premises,
message_created_at=format_datetime_utc(message_created_at),
)
documents_to_create.append(
schemas.DocumentCreate(
content=obs_content,
session_name=session_name,
level=obs_level,
metadata=metadata,
embedding=embedding,
)
)
# Use bulk creation with optional duplicate detection
accepted_documents_result = await crud.create_documents(
db,
documents_to_create,
self.workspace_name,
observer=self.observer,
observed=self.observed,
deduplicate=settings.DERIVER.DEDUPLICATE,
)
if message_level_configuration.dream.enabled:
try:
await check_and_schedule_dream(db, collection)
except Exception as e:
logger.warning(f"Failed to check dream scheduling: {e}")
return accepted_documents_result
async def get_working_representation(
self,
*,
db: AsyncSession | None = None,
session_allowlist: list[str] | None = None,
include_semantic_query: str | None = None,
embedding: list[float] | None = None,
semantic_search_top_k: int | None = None,
semantic_search_max_distance: float | None = None,
include_most_derived: bool = False,
max_observations: int = settings.DERIVER.WORKING_REPRESENTATION_MAX_OBSERVATIONS,
parent_category: str | None = None,
embedding_purpose: EmbeddingCallPurpose = EmbeddingCallPurpose.SEARCH_MEMORY,
) -> Representation:
"""
Get working representation with flexible query options.
Args:
db: Optional database session. If provided, uses it directly;
otherwise creates a new session via tracked_db.
session_allowlist: Optional session allowlist to filter by. Applied
uniformly to every query path (semantic, most-derived, and
recent). None means no session restriction; an empty list
fail-closes to an empty representation.
include_semantic_query: Query for semantic search
embedding: Pre-computed embedding for the semantic query.
semantic_search_top_k: Number of semantic results
semantic_search_max_distance: Maximum distance for semantic search
include_most_derived: Include most derived observations
max_observations: Maximum total observations to return
parent_category: Optional workflow attribution forwarded to the
fallback embedding call when the caller didn't pre-compute
an embedding (or pre-compute failed).
embedding_purpose: Embedding call_purpose tag to use on the
fallback embed when no pre-computed embedding was supplied.
Defaults to SEARCH_MEMORY; callers whose route-level
precompute uses a more specific purpose (e.g.
SESSION_CONTEXT_SEARCH) should pass that here so the
fallback path lands in the same analytics bucket.
Returns:
Representation combining various query strategies
"""
if include_semantic_query and embedding is None:
# Best-effort precompute when caller didn't supply one (or their
# precompute was suppressed). The purpose is parameterized so
# this fallback shows up in the same telemetry bucket as the
# successful path — see embedding_purpose docstring above.
with (
suppress(Exception),
embedding_call_purpose(
embedding_purpose.value,
workspace_name=self.workspace_name,
parent_category=parent_category,
),
):
embedding = await embedding_client.embed(include_semantic_query)
if db is not None:
return await self._get_working_representation_internal(
db,
session_allowlist=session_allowlist,
include_semantic_query=include_semantic_query,
embedding=embedding,
semantic_search_top_k=semantic_search_top_k,
semantic_search_max_distance=semantic_search_max_distance,
include_most_derived=include_most_derived,
max_observations=max_observations,
)
async with tracked_db(
"representation_manager.get_working_representation", read_only=True
) as new_db:
return await self._get_working_representation_internal(
new_db,
session_allowlist=session_allowlist,
include_semantic_query=include_semantic_query,
embedding=embedding,
semantic_search_top_k=semantic_search_top_k,
semantic_search_max_distance=semantic_search_max_distance,
include_most_derived=include_most_derived,
max_observations=max_observations,
)
# Private helper methods
async def _get_working_representation_internal(
self,
db: AsyncSession,
*,
session_allowlist: list[str] | None = None,
include_semantic_query: str | None = None,
embedding: list[float] | None = None,
semantic_search_top_k: int | None = None,
semantic_search_max_distance: float | None = None,
include_most_derived: bool = False,
max_observations: int = settings.DERIVER.WORKING_REPRESENTATION_MAX_OBSERVATIONS,
) -> Representation:
"""Internal implementation of get_working_representation."""
# Fail closed on an empty allowlist. This must short-circuit before
# any query: downstream stores drop an `IN ()` clause with an empty
# list (lancedb), which would silently widen the scope instead.
if session_allowlist is not None and not session_allowlist:
return Representation()
total = max_observations
# Calculate how many observations to get from each source.
# Floor of 1 when a semantic query was explicitly requested.
semantic_observations = (
min(
max(
1,
semantic_search_top_k
if semantic_search_top_k is not None
else total // 3,
),
total,
)
if include_semantic_query
else 0
)
if include_semantic_query and include_most_derived:
# three-way blend: both semantic and derived requested
top_observations = min(max(0, total // 3), total - semantic_observations)
elif include_most_derived:
# two-way blend: only derived requested
top_observations = min(max(0, total // 2), total - semantic_observations)
else:
# no derived observations requested
top_observations = 0
# remaining observations are recent
recent_observations = total - semantic_observations - top_observations
representation = Representation()
# Get semantic observations if requested
if include_semantic_query:
semantic_docs = await self._query_documents_semantic(
db,
query=include_semantic_query,
top_k=semantic_observations,
max_distance=semantic_search_max_distance,
embedding=embedding,
session_allowlist=session_allowlist,
)
representation.merge_representation(
Representation.from_documents(semantic_docs)
)
# Get most derived observations if requested
if include_most_derived:
derived_docs = await self._query_documents_most_derived(
db, top_k=top_observations, session_allowlist=session_allowlist
)
representation.merge_representation(
Representation.from_documents(derived_docs)
)
# Get recent observations
recent_docs = await self._query_documents_recent(
db, top_k=recent_observations, session_allowlist=session_allowlist
)
representation.merge_representation(Representation.from_documents(recent_docs))
return representation
async def _query_documents_semantic(
self,
db: AsyncSession,
query: str,
top_k: int,
max_distance: float | None = None,
level: str | None = None,
embedding: list[float] | None = None,
session_allowlist: list[str] | None = None,
) -> list[models.Document]:
"""Query documents by semantic similarity."""
try:
if level:
return await self._query_documents_for_level(
db,
query,
level,
top_k,
max_distance,
embedding=embedding,
session_allowlist=session_allowlist,
)
else:
documents = await crud.query_documents(
db,
workspace_name=self.workspace_name,
observer=self.observer,
observed=self.observed,
query=query,
max_distance=max_distance,
top_k=top_k,
embedding=embedding,
filters=self._build_filter_conditions(
session_allowlist=session_allowlist
)
or None,
)
db.expunge_all()
return list(documents)
except Exception as e:
logger.error(f"Error getting relevant observations: {e}")
return []
async def _query_documents_recent(
self, db: AsyncSession, top_k: int, session_allowlist: list[str] | None = None
) -> list[models.Document]:
"""Query most recent documents."""
stmt = (
select(models.Document)
.limit(top_k)
.where(
models.Document.workspace_name == self.workspace_name,
models.Document.observer == self.observer,
models.Document.observed == self.observed,
models.Document.deleted_at.is_(None),
*(
[
models.Document.session_name.in_(session_allowlist),
# Only levels with a trustworthy session stamp are
# scopeable — see ALLOWLIST_SAFE_LEVELS.
models.Document.level.in_(ALLOWLIST_SAFE_LEVELS),
]
if session_allowlist is not None
else []
),
)
.order_by(models.Document.created_at.desc())
)
result = await db.execute(stmt)
documents = result.scalars().all()
db.expunge_all()
return list(documents)
async def _query_documents_most_derived(
self, db: AsyncSession, top_k: int, session_allowlist: list[str] | None = None
) -> list[models.Document]:
"""Query most derived documents."""
stmt = (
select(models.Document)
.limit(top_k)
.where(
models.Document.workspace_name == self.workspace_name,
models.Document.observer == self.observer,
models.Document.observed == self.observed,
models.Document.deleted_at.is_(None),
*(
[
models.Document.session_name.in_(session_allowlist),
# Only levels with a trustworthy session stamp are
# scopeable — see ALLOWLIST_SAFE_LEVELS.
models.Document.level.in_(ALLOWLIST_SAFE_LEVELS),
]
if session_allowlist is not None
else []
),
)
.order_by(
models.Document.times_derived.desc(),
models.Document.created_at.desc(),
# created_at is the transaction timestamp, so documents created
# in the same batch share it -- id keeps the order deterministic.
models.Document.id,
)
)
result = await db.execute(stmt)
documents = result.scalars().all()
db.expunge_all()
return list(documents)
async def _get_observations_internal(
self,
db: AsyncSession,
query: str,
top_k: int,
max_distance: float,
level: str | None,
) -> list[models.Document]:
"""Internal method that does the actual observation retrieval."""
return await self._query_documents_semantic(
db, query, top_k, max_distance, level
)
async def _query_documents_for_level(
self,
db: AsyncSession,
query: str,
level: str,
count: int,
max_distance: float | None = None,
embedding: list[float] | None = None,
session_allowlist: list[str] | None = None,
) -> list[models.Document]:
"""Query documents for a specific level."""
documents = await crud.query_documents(
db,
workspace_name=self.workspace_name,
observer=self.observer,
observed=self.observed,
query=query,
max_distance=max_distance,
top_k=count,
filters=self._build_filter_conditions(
level, session_allowlist=session_allowlist
),
embedding=embedding,
)
# Sort by creation time
docs_sorted: list[models.Document] = sorted(
list(documents), key=lambda x: x.created_at, reverse=True
)
return docs_sorted
def _build_filter_conditions(
self,
level: str | None = None,
session_allowlist: list[str] | None = None,
) -> dict[str, Any]:
"""
Build filter conditions for document queries.
Returns a flat dict of key-value pairs for vector store filtering.
Callers must not pass an empty session_allowlist list — empty allowlists
fail closed before any query is issued (see
_get_working_representation_internal).
"""
filters: dict[str, Any] = {}
if level:
filters["level"] = level
# `is not None` (not truthiness): an explicit empty allowlist must emit
# an empty `in` so downstream stores fail closed, matching
# _query_documents_recent / _query_documents_most_derived. Truthiness
# here would silently drop the filter and widen scope.
if session_allowlist is not None:
filters["session_name"] = {"in": session_allowlist}
# Only levels with a trustworthy session stamp are scopeable. This
# overrides any narrower `level` above; an empty intersection emits
# `{"in": []}`, which matches nothing rather than everything.
filters["level"] = {"in": allowlist_safe_levels([level] if level else None)}
return filters
# Module-level functions for backward compatibility and convenience
async def get_working_representation(
workspace_name: str,
*,
db: AsyncSession | None = None,
observer: str,
observed: str,
session_allowlist: list[str] | None = None,
include_semantic_query: str | None = None,
embedding: list[float] | None = None,
semantic_search_top_k: int | None = None,
semantic_search_max_distance: float | None = None,
include_most_derived: bool = False,
max_observations: int = settings.DERIVER.WORKING_REPRESENTATION_MAX_OBSERVATIONS,
parent_category: str | None = None,
embedding_purpose: EmbeddingCallPurpose = EmbeddingCallPurpose.SEARCH_MEMORY,
) -> Representation:
"""
Get raw working representation data from the relevant document collection.
This is a convenience function that creates a RepresentationManager and calls
get_working_representation on it.
Args:
db: Optional database session. If provided, uses it directly;
otherwise creates a new session via tracked_db.
embedding: Pre-computed embedding for the semantic query.
parent_category: Workflow attribution forwarded to the fallback
embedding call when no pre-computed embedding was supplied.
embedding_purpose: Embedding call_purpose for the fallback embed;
callers should match it to whatever purpose their route-level
precompute used so failure/retry paths stay in the same bucket.
"""
manager = RepresentationManager(
workspace_name=workspace_name,
observer=observer,
observed=observed,
)
return await manager.get_working_representation(
db=db,
session_allowlist=session_allowlist,
include_semantic_query=include_semantic_query,
embedding=embedding,
semantic_search_top_k=semantic_search_top_k,
semantic_search_max_distance=semantic_search_max_distance,
include_most_derived=include_most_derived,
max_observations=max_observations,
parent_category=parent_category,
embedding_purpose=embedding_purpose,
)