Files
duocthu/apps/ai-service/rag/instrumentation.py
T

241 lines
8.0 KiB
Python

"""Non-invasive instrumentation wrappers for the live RAG object graph.
Claude owns ``rag/agent.py`` and ``rag/answer.py`` in the shared worktree.
Subclasses here add spans at their stable method boundaries without changing
those files or duplicating their domain decisions.
"""
from __future__ import annotations
from typing import Any
from .agent import RagAgent
from .answer import GroundedAnswerService
from .metrics import CLARIFY_ASKED, Metrics, PROVIDER_FAILURE, RETRIEVAL_ROUTE
from .service import RetrievalService
from .telemetry import (
annotate_current_span,
current_stage,
dependency_span,
stage,
)
def _provider_name(delegate: Any) -> str:
name = delegate.__class__.__name__.casefold()
if "cohere" in name:
return "bedrock_cohere"
if "converse" in name:
return "bedrock_converse"
if "claude" in name:
return "bedrock_claude"
return "other"
def _failure_reason(exc: BaseException) -> str:
name = exc.__class__.__name__.casefold()
if "timeout" in name:
return "timeout"
if "budget" in name:
return "request_budget_exhausted"
if "unavailable" in name or "connection" in name:
return "provider_unavailable"
return "error"
class InstrumentedGenerator:
def __init__(self, delegate: Any, metrics: Metrics) -> None:
self._delegate = delegate
self._metrics = metrics
self._provider = _provider_name(delegate)
@property
def model_id(self) -> str:
return self._delegate.model_id
def generate(self, system: str, user: str, schema: dict) -> str:
operation = {
"understanding": "understand",
"generation": "generate",
"entailment": "entailment",
}.get(current_stage(), "generate")
with dependency_span(self._provider, operation):
try:
return self._delegate.generate(system, user, schema)
except Exception as exc:
self._metrics.increment(
PROVIDER_FAILURE,
provider=self._provider,
operation=operation,
reason=_failure_reason(exc),
)
raise
class InstrumentedEmbedder:
def __init__(self, delegate: Any, metrics: Metrics) -> None:
self._delegate = delegate
self._metrics = metrics
@property
def dimensions(self) -> int:
return self._delegate.dimensions
@property
def model_id(self) -> str:
return self._delegate.model_id
def embed_query(self, text: str) -> list[float]:
with dependency_span("bedrock_cohere", "embed"):
try:
return self._delegate.embed_query(text)
except Exception as exc:
self._metrics.increment(
PROVIDER_FAILURE,
provider="bedrock_cohere",
operation="embed",
reason=_failure_reason(exc),
)
raise
class InstrumentedReranker:
def __init__(self, delegate: Any, metrics: Metrics) -> None:
self._delegate = delegate
self._metrics = metrics
def rerank(
self, query: str, documents: list[str], top_n: int | None = None
) -> list[int]:
with dependency_span("bedrock_cohere", "rerank"):
try:
return self._delegate.rerank(query, documents, top_n=top_n)
except Exception as exc:
self._metrics.increment(
PROVIDER_FAILURE,
provider="bedrock_cohere",
operation="rerank",
reason=_failure_reason(exc),
)
raise
class InstrumentedRetrievalService(RetrievalService):
def __init__(self, *args, metrics: Metrics, **kwargs) -> None:
super().__init__(*args, **kwargs)
self._observability_metrics = metrics
def retrieve_framed(self, *args, **kwargs):
with stage("retrieval"):
section_key = args[1] if len(args) > 1 else kwargs.get("section_key")
is_overview = args[3] if len(args) > 3 else kwargs.get("is_overview", False)
route = "section" if section_key else (
"overview" if is_overview else "similarity"
)
self._observability_metrics.increment(RETRIEVAL_ROUTE, route=route)
try:
result = super().retrieve_framed(*args, **kwargs)
except Exception as exc:
self._record_retrieval_failure(exc)
raise
self._annotate_result(result)
return result
def retrieve(self, *args, **kwargs):
with stage("retrieval"):
self._observability_metrics.increment(RETRIEVAL_ROUTE, route="other")
try:
result = super().retrieve(*args, **kwargs)
except Exception as exc:
self._record_retrieval_failure(exc)
raise
self._annotate_result(result)
return result
def retrieve_by_indication(self, *args, **kwargs):
with stage("retrieval"):
self._observability_metrics.increment(RETRIEVAL_ROUTE, route="indication")
try:
result = super().retrieve_by_indication(*args, **kwargs)
except Exception as exc:
self._record_retrieval_failure(exc)
raise
self._annotate_result(result)
return result
def _record_retrieval_failure(self, exc: BaseException) -> None:
self._observability_metrics.increment(
PROVIDER_FAILURE,
provider="qdrant",
operation="retrieve",
reason=_failure_reason(exc),
)
@staticmethod
def _annotate_result(result) -> None:
annotate_current_span(
**{
"duocthu.decision": result.decision.value,
"duocthu.reason": result.reason,
"duocthu.evidence_count": len(result.evidence),
}
)
def _rerank(self, query, hits):
with stage("rerank", configured=self._reranker is not None):
return super()._rerank(query, hits)
def _decide(self, evidence, is_drug_overview=False):
with stage("evidence", evidence_count=len(evidence)):
return super()._decide(evidence, is_drug_overview=is_drug_overview)
class InstrumentedGroundedAnswerService(GroundedAnswerService):
def _generate(self, *args, **kwargs):
with stage("generation"):
return super()._generate(*args, **kwargs)
def _verify_entailment(self, *args, **kwargs):
with stage("entailment"):
return super()._verify_entailment(*args, **kwargs)
class InstrumentedQueryUnderstander:
def __init__(self, delegate: Any) -> None:
self._delegate = delegate
def understand(self, *args, **kwargs):
with stage("understanding"):
frame = self._delegate.understand(*args, **kwargs)
annotate_current_span(
**{
"duocthu.turn_type": frame.turn_type,
"duocthu.needs_clarify": frame.needs_clarify,
"duocthu.system_error": frame.system_error,
}
)
return frame
class InstrumentedRagAgent(RagAgent):
def __init__(self, *args, metrics: Metrics, **kwargs) -> None:
super().__init__(*args, **kwargs)
self._observability_metrics = metrics
def handle(self, *args, **kwargs):
reply = super().handle(*args, **kwargs)
if reply.decision == "clarify":
self._observability_metrics.increment(CLARIFY_ASKED, reason=reply.reason)
return reply
def _get_history(self, *args, **kwargs):
with stage("context"):
return super()._get_history(*args, **kwargs)
def _route(self, *args, **kwargs):
with stage("routing"):
return super()._route(*args, **kwargs)
def _remember(self, *args, **kwargs):
with stage("persistence"):
return super()._remember(*args, **kwargs)