"""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)