Ajustes conforme relatorio de testes 2026-08-27
This commit is contained in:
@@ -0,0 +1,4 @@
|
||||
__all__ = ['settings']
|
||||
from .config.settings import settings
|
||||
|
||||
from .idempotency import IdempotencyStore, InMemoryIdempotencyStore, create_idempotency_store
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,12 @@
|
||||
from .publisher import AnalyticsPublisher, NoopAnalyticsPublisher
|
||||
from .composite_publisher import CompositeAnalyticsPublisher
|
||||
from .event_builder import build_analytics_event
|
||||
from .factory import create_analytics_publisher
|
||||
|
||||
__all__ = [
|
||||
"AnalyticsPublisher",
|
||||
"NoopAnalyticsPublisher",
|
||||
"CompositeAnalyticsPublisher",
|
||||
"build_analytics_event",
|
||||
"create_analytics_publisher",
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,35 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Any, Iterable
|
||||
|
||||
from .publisher import AnalyticsPublisher
|
||||
|
||||
logger = logging.getLogger("agent_framework.analytics.composite")
|
||||
|
||||
|
||||
class CompositeAnalyticsPublisher(AnalyticsPublisher):
|
||||
"""Publica o mesmo evento em múltiplos destinos.
|
||||
|
||||
Use para rodar OCI Streaming e Pub/Sub em paralelo durante transição,
|
||||
homologação ou estratégia multi-cloud.
|
||||
"""
|
||||
|
||||
def __init__(self, publishers: Iterable[AnalyticsPublisher], *, fail_silent: bool = True):
|
||||
self.publishers = list(publishers)
|
||||
self.fail_silent = fail_silent
|
||||
|
||||
async def publish(self, event_type: str, payload: dict[str, Any]) -> None:
|
||||
if not self.publishers:
|
||||
return
|
||||
|
||||
async def _safe_publish(publisher: AnalyticsPublisher) -> None:
|
||||
try:
|
||||
await publisher.publish(event_type, payload)
|
||||
except Exception:
|
||||
logger.exception("analytics.publisher_failed provider=%s event_type=%s", publisher.__class__.__name__, event_type)
|
||||
if not self.fail_silent:
|
||||
raise
|
||||
|
||||
await asyncio.gather(*[_safe_publish(p) for p in self.publishers])
|
||||
@@ -0,0 +1,27 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
|
||||
def build_analytics_event(
|
||||
event_type: str,
|
||||
payload: dict[str, Any] | None = None,
|
||||
*,
|
||||
source: str = "agent_framework",
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Monta envelope uniforme para IC/NOC/GRL.
|
||||
|
||||
O campo metadata.noc=true é preservado para que o Observer consiga rotear
|
||||
eventos também para NOC/OTEL/Elastic quando aplicável.
|
||||
"""
|
||||
body = dict(payload or {})
|
||||
meta = dict(metadata or {})
|
||||
return {
|
||||
"eventType": event_type,
|
||||
"source": source,
|
||||
"eventDate": datetime.now(timezone.utc).isoformat(),
|
||||
"payload": body,
|
||||
"metadata": meta,
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from .composite_publisher import CompositeAnalyticsPublisher
|
||||
from .publisher import AnalyticsPublisher, NoopAnalyticsPublisher
|
||||
|
||||
logger = logging.getLogger("agent_framework.analytics.factory")
|
||||
|
||||
|
||||
def _split_csv(value: str | None) -> list[str]:
|
||||
return [item.strip().lower() for item in (value or "").split(",") if item.strip()]
|
||||
|
||||
|
||||
def create_analytics_publisher(settings: Any | None = None) -> AnalyticsPublisher:
|
||||
"""Cria publisher conforme env/config.
|
||||
|
||||
Variáveis novas compatíveis:
|
||||
- ENABLE_ANALYTICS=true|false
|
||||
- ANALYTICS_PROVIDERS=oci_streaming,pubsub
|
||||
- GCP_PUBSUB_TOPIC_PATH=projects/.../topics/...
|
||||
- AGENT_PUBSUB_TOPIC=projects/.../topics/... # compatibilidade FIRST/TIM
|
||||
- GCP_PROJECT_ID=... + GCP_PUBSUB_TOPIC=...
|
||||
"""
|
||||
if settings is None:
|
||||
from agent_framework.config.settings import settings as default_settings
|
||||
settings = default_settings
|
||||
|
||||
analytics_enabled = bool(getattr(settings, "ENABLE_ANALYTICS", False))
|
||||
langfuse_enabled = bool(getattr(settings, "ENABLE_LANGFUSE", False))
|
||||
|
||||
# Historicamente o observer era usado para enviar IC/NOC/GRL ao Langfuse
|
||||
# mesmo quando o pipeline de analytics/streaming não estava habilitado.
|
||||
# Portanto, ENABLE_LANGFUSE=true também ativa o publisher Langfuse do observer.
|
||||
if not analytics_enabled and not langfuse_enabled:
|
||||
return NoopAnalyticsPublisher()
|
||||
|
||||
providers = _split_csv(getattr(settings, "ANALYTICS_PROVIDERS", "")) or ["oci_streaming"]
|
||||
if langfuse_enabled and "langfuse" not in providers:
|
||||
providers.insert(0, "langfuse")
|
||||
|
||||
# Se analytics geral estiver desligado, publica somente no Langfuse para
|
||||
# evitar inicializar OCI Streaming/PubSub por engano em ambientes locais.
|
||||
if not analytics_enabled:
|
||||
providers = [p for p in providers if p in {"langfuse", "noop", "none"}] or ["langfuse"]
|
||||
publishers: list[AnalyticsPublisher] = []
|
||||
|
||||
for provider in providers:
|
||||
try:
|
||||
if provider == "langfuse":
|
||||
from .providers.langfuse import LangfuseAnalyticsPublisher
|
||||
publishers.append(LangfuseAnalyticsPublisher(settings=settings))
|
||||
elif provider == "oci_streaming":
|
||||
from .providers.oci_streaming import OCIStreamingAnalyticsPublisher
|
||||
publishers.append(OCIStreamingAnalyticsPublisher(settings=settings))
|
||||
elif provider in {"pubsub", "gcp_pubsub", "gcp"}:
|
||||
from .providers.pubsub import PubSubAnalyticsPublisher
|
||||
topic = (
|
||||
getattr(settings, "GCP_PUBSUB_TOPIC_PATH", None)
|
||||
or getattr(settings, "AGENT_PUBSUB_TOPIC", None)
|
||||
)
|
||||
publishers.append(PubSubAnalyticsPublisher(topic_path=topic))
|
||||
elif provider in {"noop", "none"}:
|
||||
publishers.append(NoopAnalyticsPublisher())
|
||||
else:
|
||||
logger.warning("analytics.provider_ignored provider=%s", provider)
|
||||
except Exception:
|
||||
logger.exception("analytics.provider_init_failed provider=%s", provider)
|
||||
|
||||
if not publishers:
|
||||
# Sem este log, "analytics ligado mas todos os providers falharam" fica
|
||||
# indistinguivel de "analytics desligado": o publisher no-op descarta
|
||||
# IC/NOC/GRL em silencio ate o processo ser reiniciado.
|
||||
logger.error(
|
||||
"analytics.no_publisher_available providers=%s enable_analytics=%s "
|
||||
"enable_langfuse=%s; telemetria sera descartada ate o proximo restart",
|
||||
",".join(providers),
|
||||
analytics_enabled,
|
||||
langfuse_enabled,
|
||||
)
|
||||
return NoopAnalyticsPublisher()
|
||||
if len(publishers) == 1:
|
||||
return publishers[0]
|
||||
return CompositeAnalyticsPublisher(publishers)
|
||||
@@ -0,0 +1,11 @@
|
||||
from .oci_streaming import OCIStreamingAnalyticsPublisher
|
||||
from .pubsub import PubSubAnalyticsPublisher
|
||||
from .kafka import KafkaAnalyticsPublisher
|
||||
from .langfuse import LangfuseAnalyticsPublisher
|
||||
|
||||
__all__ = [
|
||||
"OCIStreamingAnalyticsPublisher",
|
||||
"PubSubAnalyticsPublisher",
|
||||
"KafkaAnalyticsPublisher",
|
||||
"LangfuseAnalyticsPublisher",
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,25 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from agent_framework.analytics.publisher import AnalyticsPublisher
|
||||
|
||||
|
||||
class KafkaAnalyticsPublisher(AnalyticsPublisher):
|
||||
"""Publisher Kafka opcional.
|
||||
|
||||
Recebe um producer já criado para não acoplar o framework a uma lib específica
|
||||
(confluent-kafka, aiokafka, kafka-python etc.). O producer precisa expor send
|
||||
assíncrono ou síncrono.
|
||||
"""
|
||||
|
||||
def __init__(self, producer: Any, topic: str):
|
||||
self.producer = producer
|
||||
self.topic = topic
|
||||
|
||||
async def publish(self, event_type: str, payload: dict[str, Any]) -> None:
|
||||
message = json.dumps({"type": event_type, "payload": payload}, default=str).encode("utf-8")
|
||||
result = self.producer.send(self.topic, key=event_type.encode("utf-8"), value=message)
|
||||
if hasattr(result, "__await__"):
|
||||
await result
|
||||
@@ -0,0 +1,446 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from agent_framework.analytics.publisher import AnalyticsPublisher
|
||||
from agent_framework.observability.code_mapper import create_observability_code_mapper
|
||||
|
||||
try: # Avoid making analytics import fragile in old deployments.
|
||||
from agent_framework.observability.context import get_current_observation_id, get_observability_context
|
||||
except Exception: # pragma: no cover
|
||||
get_observability_context = None # type: ignore
|
||||
get_current_observation_id = None # type: ignore
|
||||
|
||||
logger = logging.getLogger("agent_framework.analytics.langfuse")
|
||||
|
||||
|
||||
def _truthy(value: Any, default: bool = False) -> bool:
|
||||
if value is None:
|
||||
return default
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
return str(value).strip().lower() in {"1", "true", "yes", "on", "y"}
|
||||
|
||||
|
||||
def _safe_metadata(value: Any) -> Any:
|
||||
"""Remove/mascara segredos antes de enviar metadata para Langfuse."""
|
||||
if isinstance(value, dict):
|
||||
out: dict[str, Any] = {}
|
||||
for key, item in value.items():
|
||||
lk = str(key).lower()
|
||||
if any(token in lk for token in ("password", "secret", "token", "api_key", "authorization")):
|
||||
out[key] = "***"
|
||||
else:
|
||||
out[key] = _safe_metadata(item)
|
||||
return out
|
||||
if isinstance(value, list):
|
||||
return [_safe_metadata(item) for item in value]
|
||||
return value
|
||||
|
||||
|
||||
_LANGFUSE_TRACE_ID_RE = re.compile(r"^[0-9a-f]{32}$")
|
||||
_INTERNAL_PREFIXES = ("IC.", "AGA.", "NOC.", "GRL.")
|
||||
_TECHNICAL_PREFIXES = (
|
||||
"langgraph.",
|
||||
"mcp.",
|
||||
"guardrail.",
|
||||
"judge.",
|
||||
"workflow.",
|
||||
"rag.",
|
||||
"cache.",
|
||||
"checkpoint.",
|
||||
)
|
||||
|
||||
|
||||
def _clean_str(value: Any) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
text = str(value).strip()
|
||||
return text or None
|
||||
|
||||
|
||||
def _first(*values: Any) -> str | None:
|
||||
for value in values:
|
||||
text = _clean_str(value)
|
||||
if text:
|
||||
return text
|
||||
return None
|
||||
|
||||
|
||||
def _current_context() -> dict[str, Any]:
|
||||
if get_observability_context is None:
|
||||
return {}
|
||||
try:
|
||||
return get_observability_context().clean()
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def _current_parent_observation_id() -> str | None:
|
||||
if get_current_observation_id is None:
|
||||
return None
|
||||
try:
|
||||
value = get_current_observation_id()
|
||||
return str(value) if value else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _is_internal_name(name: Any) -> bool:
|
||||
text = _clean_str(name) or ""
|
||||
return text.startswith(_INTERNAL_PREFIXES)
|
||||
|
||||
|
||||
def _is_technical_name(name: Any) -> bool:
|
||||
text = _clean_str(name) or ""
|
||||
return text.startswith(_TECHNICAL_PREFIXES)
|
||||
|
||||
|
||||
def _is_control_or_technical(name: Any) -> bool:
|
||||
return _is_internal_name(name) or _is_technical_name(name)
|
||||
|
||||
|
||||
def _extract_envelope_event_type(envelope: dict[str, Any]) -> str | None:
|
||||
return _first(
|
||||
envelope.get("eventType"),
|
||||
envelope.get("event_type"),
|
||||
envelope.get("name"),
|
||||
envelope.get("type"),
|
||||
)
|
||||
|
||||
|
||||
def _is_wrapped_internal_event(event_type: str, envelope: dict[str, Any]) -> bool:
|
||||
"""Detecta caso que gerava trace raiz errado.
|
||||
|
||||
Exemplo observado no Langfuse:
|
||||
name=http.request.completed
|
||||
input={"eventType": "NOC.006", ...}
|
||||
output={"published": true}
|
||||
|
||||
Isso não é o trace real da request; é apenas o publisher de analytics
|
||||
emitindo um envelope IC/NOC/GRL através de um evento técnico. Esse registro
|
||||
deve ser suprimido para não poluir a tela Tracing -> Traces.
|
||||
"""
|
||||
envelope_event_type = _extract_envelope_event_type(envelope)
|
||||
return bool(
|
||||
envelope_event_type
|
||||
and _is_internal_name(envelope_event_type)
|
||||
and str(event_type) != envelope_event_type
|
||||
and str(event_type).startswith(("http.request.", "gateway.", "telemetry."))
|
||||
)
|
||||
|
||||
|
||||
def _raw_correlation_id(metadata: dict[str, Any]) -> str | None:
|
||||
# IMPORTANT: prefer request/trace ids over transaction/session ids. Using
|
||||
# transaction/session as first choice created duplicate root traces for
|
||||
# IC/NOC/GRL events while the HTTP trace used request_id.
|
||||
value = (
|
||||
metadata.get("traceId")
|
||||
or metadata.get("trace_id")
|
||||
or metadata.get("requestId")
|
||||
or metadata.get("request_id")
|
||||
or metadata.get("transactionId")
|
||||
or metadata.get("transaction_id")
|
||||
or metadata.get("sessionId")
|
||||
or metadata.get("session_id")
|
||||
)
|
||||
return str(value) if value else None
|
||||
|
||||
|
||||
def _langfuse_trace_id(value: Any) -> str | None:
|
||||
"""Normaliza ids do framework/business para o formato aceito pelo Langfuse.
|
||||
|
||||
Langfuse SDK v3 exige 32 caracteres hex minúsculos. UUIDs com hífens são
|
||||
compactados; ids de negócio/sessão viram hash md5 determinístico.
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
raw = str(value).strip().lower()
|
||||
if not raw:
|
||||
return None
|
||||
compact = raw.replace("-", "")
|
||||
if _LANGFUSE_TRACE_ID_RE.match(compact):
|
||||
return compact
|
||||
return hashlib.md5(raw.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def _correlation_trace_id(metadata: dict[str, Any]) -> str | None:
|
||||
return _langfuse_trace_id(_raw_correlation_id(metadata))
|
||||
|
||||
|
||||
def _with_trace_context(kwargs: dict[str, Any], metadata: dict[str, Any]) -> dict[str, Any]:
|
||||
raw_id = _raw_correlation_id(metadata)
|
||||
trace_id = _langfuse_trace_id(raw_id)
|
||||
parent_id = (
|
||||
metadata.get("parent_observation_id")
|
||||
or metadata.get("parent_span_id")
|
||||
or kwargs.get("parent_observation_id")
|
||||
or kwargs.get("parent_span_id")
|
||||
or _current_parent_observation_id()
|
||||
)
|
||||
if trace_id:
|
||||
trace_context = dict(kwargs.get("trace_context") or {})
|
||||
trace_context.setdefault("trace_id", trace_id)
|
||||
if parent_id:
|
||||
trace_context.setdefault("parent_span_id", str(parent_id))
|
||||
kwargs["trace_context"] = trace_context
|
||||
meta = kwargs.setdefault("metadata", {})
|
||||
if isinstance(meta, dict):
|
||||
meta.setdefault("framework_trace_id", raw_id)
|
||||
meta.setdefault("langfuse_trace_id", trace_id)
|
||||
if parent_id:
|
||||
meta.setdefault("parent_observation_id", str(parent_id))
|
||||
return kwargs
|
||||
|
||||
|
||||
def _allow_standalone_internal_events() -> bool:
|
||||
# Default false: IC/NOC/GRL sem contexto de request não devem criar linhas
|
||||
# soltas na tela principal de Traces. Habilite só para debug isolado.
|
||||
return _truthy(os.getenv("LANGFUSE_ALLOW_STANDALONE_INTERNAL_EVENTS"), False)
|
||||
|
||||
|
||||
class LangfuseAnalyticsPublisher(AnalyticsPublisher):
|
||||
"""Publica eventos IC/NOC/GRL no Langfuse sem criar traces raiz duplicados.
|
||||
|
||||
Regra principal:
|
||||
- 1 request/workflow = 1 trace raiz;
|
||||
- IC/NOC/GRL e eventos técnicos entram como observations/spans dentro do
|
||||
trace corrente;
|
||||
- envelopes internos embrulhados em eventos HTTP/gateway não criam trace
|
||||
próprio com output {"published": true}.
|
||||
"""
|
||||
|
||||
def __init__(self, settings: Any | None = None, langfuse: Any | None = None):
|
||||
if settings is None:
|
||||
from agent_framework.config.settings import settings as default_settings
|
||||
settings = default_settings
|
||||
|
||||
self.settings = settings
|
||||
self.code_mapper = create_observability_code_mapper(settings)
|
||||
self.langfuse = langfuse
|
||||
self.enabled = True
|
||||
|
||||
if self.langfuse is not None:
|
||||
return
|
||||
|
||||
public_key = getattr(settings, "LANGFUSE_PUBLIC_KEY", None) or os.getenv("LANGFUSE_PUBLIC_KEY")
|
||||
secret_key = getattr(settings, "LANGFUSE_SECRET_KEY", None) or os.getenv("LANGFUSE_SECRET_KEY")
|
||||
host = getattr(settings, "LANGFUSE_HOST", None) or os.getenv("LANGFUSE_HOST") or "https://cloud.langfuse.com"
|
||||
|
||||
if not public_key or not secret_key:
|
||||
self.enabled = False
|
||||
logger.warning("LangfuseAnalyticsPublisher desabilitado: LANGFUSE_PUBLIC_KEY/LANGFUSE_SECRET_KEY ausentes")
|
||||
return
|
||||
|
||||
try:
|
||||
from langfuse import Langfuse # type: ignore
|
||||
self.langfuse = Langfuse(public_key=public_key, secret_key=secret_key, host=host)
|
||||
logger.info("LangfuseAnalyticsPublisher habilitado host=%s", host)
|
||||
except Exception:
|
||||
self.enabled = False
|
||||
self.langfuse = None
|
||||
logger.exception("Falha ao inicializar LangfuseAnalyticsPublisher")
|
||||
|
||||
async def publish(self, event_type: str, payload: dict[str, Any]) -> None:
|
||||
if not self.enabled or self.langfuse is None:
|
||||
return
|
||||
|
||||
event_type = str(event_type)
|
||||
envelope = dict(payload or {})
|
||||
|
||||
# Prevent the exact pollution seen in Langfuse: http.request.completed
|
||||
# traces whose input is a NOC/IC envelope and output is {published:true}.
|
||||
if _is_wrapped_internal_event(event_type, envelope):
|
||||
logger.debug(
|
||||
"langfuse.analytics.skip_wrapped_internal event_type=%s envelope_event_type=%s",
|
||||
event_type,
|
||||
_extract_envelope_event_type(envelope),
|
||||
)
|
||||
return
|
||||
|
||||
body = envelope.get("payload") if isinstance(envelope.get("payload"), dict) else {}
|
||||
metadata = envelope.get("metadata") if isinstance(envelope.get("metadata"), dict) else {}
|
||||
ctx = _current_context()
|
||||
|
||||
source = envelope.get("source") or "agent_framework"
|
||||
event_date = envelope.get("eventDate")
|
||||
envelope_event_type = _extract_envelope_event_type(envelope)
|
||||
effective_event_type = envelope_event_type if _is_internal_name(envelope_event_type) else event_type
|
||||
|
||||
# LangfuseAnalyticsPublisher talks directly to the Langfuse SDK and does
|
||||
# not pass through Telemetry._start_observation(). Apply the same contract
|
||||
# mapper here so analytics observations cannot leak internal names.
|
||||
original_effective_event_type = str(effective_event_type)
|
||||
effective_event_type, mapping_meta = self.code_mapper.normalize_name(
|
||||
original_effective_event_type,
|
||||
metadata,
|
||||
)
|
||||
if mapping_meta != metadata:
|
||||
metadata = mapping_meta
|
||||
if isinstance(envelope.get("metadata"), dict):
|
||||
envelope["metadata"] = dict(mapping_meta)
|
||||
|
||||
# Correlation priority: current ObservabilityContext > payload metadata >
|
||||
# transaction/session fallback. This keeps IC/NOC/GRL in the same HTTP trace.
|
||||
correlation_request_id = _first(
|
||||
ctx.get("request_id"),
|
||||
ctx.get("trace_id"),
|
||||
body.get("request_id"), metadata.get("request_id"),
|
||||
body.get("requestId"), metadata.get("requestId"),
|
||||
envelope.get("request_id"), envelope.get("requestId"),
|
||||
)
|
||||
correlation_trace_id = _first(
|
||||
ctx.get("trace_id"),
|
||||
ctx.get("request_id"),
|
||||
body.get("trace_id"), metadata.get("trace_id"),
|
||||
body.get("traceId"), metadata.get("traceId"),
|
||||
correlation_request_id,
|
||||
)
|
||||
correlation_session_id = _first(
|
||||
ctx.get("session_id"),
|
||||
body.get("session_id"), metadata.get("session_id"),
|
||||
body.get("sessionId"), metadata.get("sessionId"),
|
||||
body.get("transaction_id"), metadata.get("transaction_id"),
|
||||
body.get("transactionId"), metadata.get("transactionId"),
|
||||
)
|
||||
|
||||
is_internal = _is_internal_name(effective_event_type)
|
||||
is_technical = _is_technical_name(effective_event_type)
|
||||
|
||||
# IC/NOC/GRL without current/request correlation are usually emitted by
|
||||
# background/legacy publishers. Do not create standalone trace rows unless
|
||||
# explicitly requested for debugging.
|
||||
if (is_internal or is_technical) and not correlation_trace_id and not _allow_standalone_internal_events():
|
||||
logger.debug("langfuse.analytics.skip_unrelated_internal event_type=%s", effective_event_type)
|
||||
return
|
||||
|
||||
langfuse_metadata = _safe_metadata({
|
||||
"eventType": effective_event_type,
|
||||
"observability_name_internal": mapping_meta.get("observability_name_internal"),
|
||||
"observability_name_mapped": mapping_meta.get("observability_name_mapped"),
|
||||
"observability_code_mapped": mapping_meta.get("observability_code_mapped"),
|
||||
"original_event_type": original_effective_event_type if original_effective_event_type != effective_event_type else (event_type if event_type != effective_event_type else None),
|
||||
"source": source,
|
||||
"eventDate": event_date,
|
||||
"payload": body,
|
||||
"metadata": metadata,
|
||||
"ic": _is_ic(str(effective_event_type), metadata),
|
||||
"noc": _is_noc(str(effective_event_type), metadata),
|
||||
"grl": _is_grl(str(effective_event_type), metadata),
|
||||
"tag": body.get("tag") or metadata.get("tag") or effective_event_type,
|
||||
"request_id": correlation_request_id,
|
||||
"trace_id": correlation_trace_id,
|
||||
"transaction_id": body.get("transaction_id") or metadata.get("transaction_id") or body.get("transactionId") or metadata.get("transactionId"),
|
||||
"sessionId": correlation_session_id,
|
||||
"session_id": correlation_session_id,
|
||||
"messageId": body.get("messageId") or metadata.get("messageId") or body.get("message_id") or metadata.get("message_id") or ctx.get("message_id"),
|
||||
"agentId": body.get("agentId") or metadata.get("agentId") or body.get("agent_id") or metadata.get("agent_id") or ctx.get("agent_id"),
|
||||
"channelId": body.get("channelId") or metadata.get("channelId") or body.get("channel") or metadata.get("channel") or ctx.get("channel"),
|
||||
"workflow_id": body.get("workflow_id") or metadata.get("workflow_id") or ctx.get("workflow_id"),
|
||||
"tenant_id": body.get("tenant_id") or metadata.get("tenant_id") or ctx.get("tenant_id"),
|
||||
"parent_observation_id": body.get("parent_observation_id") or metadata.get("parent_observation_id") or _current_parent_observation_id(),
|
||||
})
|
||||
|
||||
# Keep correlation metadata on the trace, but do not turn every control
|
||||
# event code into a trace tag. IC/NOC/GRL are represented by the child
|
||||
# observation below; tags are not a substitute for the event span and
|
||||
# high-cardinality event-code tags make the trace harder to inspect.
|
||||
self._update_current_trace(langfuse_metadata)
|
||||
|
||||
# Prefer current/correlated observation API. For internal/technical events,
|
||||
# do not fall back to standalone span/trace APIs if this fails.
|
||||
try:
|
||||
if hasattr(self.langfuse, "start_as_current_observation"):
|
||||
kwargs = {
|
||||
"name": str(effective_event_type),
|
||||
"as_type": "span",
|
||||
"input": envelope,
|
||||
"metadata": langfuse_metadata,
|
||||
}
|
||||
# trace_context rebuilds the parent as a remote span (SDK cross-process
|
||||
# propagation); skip it when a real span is already active locally.
|
||||
if not _current_parent_observation_id():
|
||||
kwargs = _with_trace_context(kwargs, langfuse_metadata)
|
||||
try:
|
||||
cm = self.langfuse.start_as_current_observation(**kwargs)
|
||||
except (TypeError, ValueError):
|
||||
kwargs.pop("trace_context", None)
|
||||
cm = self.langfuse.start_as_current_observation(**kwargs)
|
||||
with cm as observation:
|
||||
_update_observation(observation, output={"published": True})
|
||||
return
|
||||
except Exception:
|
||||
log = logger.warning if is_internal else logger.debug
|
||||
log("Falha ao publicar Langfuse observation para %s", effective_event_type, exc_info=True)
|
||||
if is_internal or is_technical:
|
||||
return
|
||||
|
||||
if is_internal or is_technical:
|
||||
return
|
||||
|
||||
# Legacy fallbacks only for non-internal, high-level events.
|
||||
try:
|
||||
trace_id = _correlation_trace_id(langfuse_metadata)
|
||||
if trace_id and hasattr(self.langfuse, "trace"):
|
||||
trace = self.langfuse.trace(
|
||||
id=str(trace_id),
|
||||
name=str(langfuse_metadata.get("request_id") or langfuse_metadata.get("sessionId") or "agent_framework.request"),
|
||||
session_id=langfuse_metadata.get("sessionId"),
|
||||
user_id=langfuse_metadata.get("user_id") or langfuse_metadata.get("userId"),
|
||||
metadata={k: v for k, v in langfuse_metadata.items() if v is not None},
|
||||
)
|
||||
if hasattr(trace, "span"):
|
||||
span = trace.span(name=str(effective_event_type), input=envelope, metadata=langfuse_metadata)
|
||||
if hasattr(span, "end"):
|
||||
span.end(output={"published": True})
|
||||
return
|
||||
except Exception:
|
||||
logger.debug("Falha ao publicar Langfuse span correlacionado para %s", effective_event_type, exc_info=True)
|
||||
|
||||
try:
|
||||
if hasattr(self.langfuse, "span"):
|
||||
span = self.langfuse.span(name=str(effective_event_type), input=envelope, metadata=langfuse_metadata)
|
||||
if hasattr(span, "end"):
|
||||
span.end(output={"published": True})
|
||||
return
|
||||
except Exception:
|
||||
logger.debug("Falha ao publicar Langfuse span legado para %s", effective_event_type, exc_info=True)
|
||||
|
||||
def _update_current_trace(self, metadata: dict[str, Any]) -> None:
|
||||
try:
|
||||
kwargs: dict[str, Any] = {
|
||||
"metadata": {k: v for k, v in metadata.items() if v is not None},
|
||||
}
|
||||
session_id = metadata.get("sessionId") or metadata.get("session_id")
|
||||
if session_id:
|
||||
kwargs["session_id"] = str(session_id)
|
||||
if hasattr(self.langfuse, "update_current_trace"):
|
||||
self.langfuse.update_current_trace(**kwargs)
|
||||
except Exception:
|
||||
logger.debug("Langfuse update_current_trace ignorado", exc_info=True)
|
||||
|
||||
|
||||
def _update_observation(observation: Any, **kwargs: Any) -> None:
|
||||
if observation is None:
|
||||
return
|
||||
try:
|
||||
if hasattr(observation, "update"):
|
||||
observation.update(**{k: v for k, v in kwargs.items() if v is not None})
|
||||
except Exception:
|
||||
logger.debug("Langfuse observation update ignorado", exc_info=True)
|
||||
|
||||
|
||||
def _is_noc(event_type: str, metadata: dict[str, Any]) -> bool:
|
||||
return event_type.startswith("NOC.") or _truthy(metadata.get("noc"))
|
||||
|
||||
|
||||
def _is_grl(event_type: str, metadata: dict[str, Any]) -> bool:
|
||||
return event_type.startswith("GRL.") or _truthy(metadata.get("grl"))
|
||||
|
||||
|
||||
def _is_ic(event_type: str, metadata: dict[str, Any]) -> bool:
|
||||
return event_type.startswith(("IC.", "AGA.")) or _truthy(metadata.get("ic"))
|
||||
@@ -0,0 +1,28 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from agent_framework.analytics.publisher import AnalyticsPublisher
|
||||
from agent_framework.analytics.tim_sequence import ensure_sequence_envelope
|
||||
|
||||
|
||||
class OCIStreamingAnalyticsPublisher(AnalyticsPublisher):
|
||||
"""Adapter para reutilizar o publisher OCI Streaming existente do framework."""
|
||||
|
||||
def __init__(self, settings: Any | None = None, event_publisher: Any | None = None):
|
||||
if event_publisher is not None:
|
||||
self.event_publisher = event_publisher
|
||||
else:
|
||||
from agent_framework.config.settings import settings as default_settings
|
||||
from agent_framework.events.oci_streaming import create_event_publisher
|
||||
self.event_publisher = create_event_publisher(settings or default_settings)
|
||||
|
||||
async def publish(self, event_type: str, payload: dict[str, Any]) -> None:
|
||||
# Carimba o contador de sequence no envelope antes do publish, espelhando o
|
||||
# PubSubAnalyticsPublisher. Sem isto o path OCI Streaming sai sem sequence
|
||||
# (a geração estava amarrada apenas ao Pub/Sub na migração do framework).
|
||||
# ensure_sequence_envelope não quebra observabilidade: se faltar sessionId
|
||||
# ou o backend do contador falhar, o evento segue sem o campo.
|
||||
if isinstance(payload, dict):
|
||||
payload = await ensure_sequence_envelope(payload)
|
||||
await self.event_publisher.publish(event_type, payload)
|
||||
@@ -0,0 +1,111 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from agent_framework.analytics.tim_payload_mapper import map_analytics_event_to_tim_flat_payload
|
||||
from agent_framework.analytics.tim_sequence import ensure_sequence
|
||||
|
||||
from agent_framework.analytics.publisher import AnalyticsPublisher
|
||||
|
||||
logger = logging.getLogger("agent_framework.analytics.pubsub")
|
||||
|
||||
|
||||
class PubSubAnalyticsPublisher(AnalyticsPublisher):
|
||||
"""Publisher GCP Pub/Sub real, compatível com FIRST/TIM.
|
||||
|
||||
Formas aceitas de configuração:
|
||||
|
||||
1. GCP_PUBSUB_TOPIC_PATH=projects/<project-id>/topics/<topic-id>
|
||||
2. AGENT_PUBSUB_TOPIC=projects/<project-id>/topics/<topic-id>
|
||||
3. GCP_PROJECT_ID=<project-id> + GCP_PUBSUB_TOPIC=<topic-id>
|
||||
|
||||
Credenciais seguem o padrão Google:
|
||||
GOOGLE_APPLICATION_CREDENTIALS=/secrets/service-account.json
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
topic_path: str | None = None,
|
||||
*,
|
||||
project_id: str | None = None,
|
||||
topic_id: str | None = None,
|
||||
ordering_key: str | None = None,
|
||||
timeout_seconds: float | None = None,
|
||||
):
|
||||
self.topic_path = self._resolve_topic_path(topic_path, project_id=project_id, topic_id=topic_id)
|
||||
self.ordering_key = ordering_key or os.getenv("GCP_PUBSUB_ORDERING_KEY") or ""
|
||||
self.timeout_seconds = float(timeout_seconds or os.getenv("GCP_PUBSUB_TIMEOUT_SECONDS") or 30)
|
||||
self.payload_mode = (os.getenv("PUBSUB_PAYLOAD_MODE") or os.getenv("ANALYTICS_PUBSUB_PAYLOAD_MODE") or "flat").strip().lower()
|
||||
self.exclude_noc = (os.getenv("PUBSUB_EXCLUDE_NOC") or "true").strip().lower() in {"1", "true", "yes", "y", "on"}
|
||||
self.excluded_event_types = {
|
||||
item.strip().upper()
|
||||
for item in os.getenv("PUBSUB_EXCLUDED_EVENT_TYPES", "").split(",")
|
||||
if item.strip()
|
||||
}
|
||||
|
||||
from google.cloud import pubsub_v1 # type: ignore
|
||||
|
||||
self.client = pubsub_v1.PublisherClient()
|
||||
|
||||
@staticmethod
|
||||
def _resolve_topic_path(topic_path: str | None, *, project_id: str | None, topic_id: str | None) -> str:
|
||||
explicit = (
|
||||
topic_path
|
||||
or os.getenv("GCP_PUBSUB_TOPIC_PATH")
|
||||
or os.getenv("AGENT_PUBSUB_TOPIC")
|
||||
or os.getenv("PUBSUB_TOPIC_PATH")
|
||||
)
|
||||
if explicit:
|
||||
explicit = explicit.strip()
|
||||
if explicit.startswith("projects/"):
|
||||
return explicit
|
||||
# Permite passar só o nome do tópico quando project_id estiver disponível.
|
||||
project = project_id or os.getenv("GCP_PROJECT_ID") or os.getenv("GOOGLE_CLOUD_PROJECT")
|
||||
if project:
|
||||
return f"projects/{project}/topics/{explicit}"
|
||||
raise ValueError("topic_path deve estar no formato projects/<project-id>/topics/<topic-id> quando GCP_PROJECT_ID não está definido")
|
||||
|
||||
project = project_id or os.getenv("GCP_PROJECT_ID") or os.getenv("GOOGLE_CLOUD_PROJECT")
|
||||
topic = topic_id or os.getenv("GCP_PUBSUB_TOPIC") or os.getenv("PUBSUB_TOPIC")
|
||||
if project and topic:
|
||||
return f"projects/{project}/topics/{topic}"
|
||||
|
||||
raise ValueError("Configure GCP_PUBSUB_TOPIC_PATH, AGENT_PUBSUB_TOPIC ou GCP_PROJECT_ID + GCP_PUBSUB_TOPIC")
|
||||
|
||||
async def publish(self, event_type: str, payload: dict[str, Any]) -> None:
|
||||
event_key = str(event_type).upper()
|
||||
if event_key in self.excluded_event_types:
|
||||
logger.debug("analytics.pubsub.skipped_event event_type=%s", event_type)
|
||||
return
|
||||
|
||||
metadata = payload.get("metadata") if isinstance(payload, dict) else None
|
||||
is_noc = str(event_type).startswith("NOC.") or (isinstance(metadata, dict) and metadata.get("noc") is True)
|
||||
if is_noc and self.exclude_noc:
|
||||
logger.debug("analytics.pubsub.skipped_noc event_type=%s", event_type)
|
||||
return
|
||||
|
||||
if self.payload_mode in {"legacy", "envelope", "wrapped"}:
|
||||
message = {"type": event_type, "payload": payload}
|
||||
else:
|
||||
message = map_analytics_event_to_tim_flat_payload(event_type, payload, keep_none=False)
|
||||
message = await ensure_sequence(message)
|
||||
|
||||
data = json.dumps(message, default=str, ensure_ascii=False).encode("utf-8")
|
||||
attributes = {
|
||||
"event_type": str(event_type),
|
||||
"source": str(payload.get("source") or "agent_framework"),
|
||||
}
|
||||
if is_noc:
|
||||
attributes["noc"] = "true"
|
||||
|
||||
kwargs: dict[str, Any] = dict(attributes)
|
||||
if self.ordering_key:
|
||||
kwargs["ordering_key"] = self.ordering_key
|
||||
|
||||
future = self.client.publish(self.topic_path, data=data, **kwargs)
|
||||
await asyncio.to_thread(future.result, timeout=self.timeout_seconds)
|
||||
logger.debug("analytics.pubsub.published event_type=%s topic=%s", event_type, self.topic_path)
|
||||
@@ -0,0 +1,27 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger("agent_framework.analytics")
|
||||
|
||||
|
||||
class AnalyticsPublisher(ABC):
|
||||
"""Contrato único para eventos analíticos corporativos.
|
||||
|
||||
A intenção é desacoplar o agente de OCI Streaming, GCP Pub/Sub, Kafka,
|
||||
BigQuery ou qualquer outro destino. Os agentes publicam eventos de negócio
|
||||
ou operação usando apenas este contrato.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def publish(self, event_type: str, payload: dict[str, Any]) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class NoopAnalyticsPublisher(AnalyticsPublisher):
|
||||
"""Publisher seguro para ambientes locais/testes."""
|
||||
|
||||
async def publish(self, event_type: str, payload: dict[str, Any]) -> None:
|
||||
logger.info("analytics.noop event_type=%s payload_keys=%s", event_type, sorted(payload.keys()))
|
||||
@@ -0,0 +1,152 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
|
||||
def _first(mapping: dict[str, Any], *keys: str) -> Any:
|
||||
for key in keys:
|
||||
if key in mapping and mapping.get(key) is not None:
|
||||
return mapping.get(key)
|
||||
return None
|
||||
|
||||
|
||||
def _as_list(value: Any) -> Any:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
if isinstance(value, (tuple, set)):
|
||||
return list(value)
|
||||
return [value]
|
||||
|
||||
|
||||
def _collect_agent_specific_data(metadata: dict[str, Any], body: dict[str, Any]) -> dict[str, Any] | None:
|
||||
prefixed: dict[str, Any] = {}
|
||||
for source in (metadata, body):
|
||||
for key, value in source.items():
|
||||
if key.startswith("agentSpecificData."):
|
||||
prefixed[key.removeprefix("agentSpecificData.")] = value
|
||||
if prefixed:
|
||||
return prefixed
|
||||
|
||||
direct = _first(metadata, "agentSpecificData")
|
||||
if isinstance(direct, dict):
|
||||
return dict(direct)
|
||||
if isinstance(direct, str) and direct.strip():
|
||||
try:
|
||||
parsed = json.loads(direct)
|
||||
if isinstance(parsed, dict):
|
||||
return parsed
|
||||
except (TypeError, ValueError, json.JSONDecodeError):
|
||||
pass
|
||||
direct = _first(body, "agentSpecificData")
|
||||
if isinstance(direct, dict):
|
||||
return dict(direct)
|
||||
if isinstance(direct, str) and direct.strip():
|
||||
try:
|
||||
parsed = json.loads(direct)
|
||||
if isinstance(parsed, dict):
|
||||
return parsed
|
||||
except (TypeError, ValueError, json.JSONDecodeError):
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def map_analytics_event_to_tim_flat_payload(
|
||||
event_type: str,
|
||||
event: dict[str, Any],
|
||||
*,
|
||||
keep_none: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
"""Map the framework analytics envelope to TIM's flat Pub/Sub/NOC schema.
|
||||
|
||||
The canonical fields are published at the JSON root. The only intentional
|
||||
nested object is ``agentSpecificData``.
|
||||
"""
|
||||
if not isinstance(event, dict):
|
||||
event = {}
|
||||
|
||||
body = event.get("payload") if isinstance(event.get("payload"), dict) else {}
|
||||
metadata = event.get("metadata") if isinstance(event.get("metadata"), dict) else {}
|
||||
data: dict[str, Any] = {**body, **metadata}
|
||||
|
||||
token_usage = event.get("token_usage") if isinstance(event.get("token_usage"), dict) else {}
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
# Tracking
|
||||
"eventType": event.get("eventType") or event_type,
|
||||
"traceId": _first(data, "traceId", "trace_id"),
|
||||
"transactionId": _first(data, "transactionId", "transaction_id", "transactionID"),
|
||||
"spanId": _first(data, "spanId", "span_id"),
|
||||
"parentSpanId": _first(data, "parentSpanId", "parent_span_id"),
|
||||
"eventName": _first(data, "eventName", "name"),
|
||||
"version": _first(data, "version") or "1.0",
|
||||
"eventDate": _first(data, "eventDate") or event.get("eventDate") or datetime.now(timezone.utc).isoformat(),
|
||||
# Session/channel
|
||||
"sessionId": _first(data, "sessionId", "session_id"),
|
||||
"channelId": _first(data, "channelId", "channel", "channel_id"),
|
||||
"agentId": _first(data, "agentId", "agent_id"),
|
||||
"customerCode": _first(data, "customerCode", "customer_code"),
|
||||
"touchpoint": _first(data, "touchpoint"),
|
||||
"protocol": _first(data, "protocol"),
|
||||
"tag": _first(data, "tag") or event.get("eventType") or event_type,
|
||||
"noc": True if _first(data, "noc") is True else None,
|
||||
# Protocol/session
|
||||
"agentProtocolId": _first(data, "agentProtocolId", "agent_protocol_id"),
|
||||
"adjustedProtocol": _first(data, "adjustedProtocol", "adjusted_protocol"),
|
||||
"sessionCreatedAt": _first(data, "sessionCreatedAt", "session_created_at"),
|
||||
"sessionEndAt": _first(data, "sessionEndAt", "session_end_at"),
|
||||
# URA/voice
|
||||
"uraCallId": _first(data, "uraCallId", "ura_call_id"),
|
||||
"transcriptionId": _first(data, "transcriptionId", "transcription_id"),
|
||||
"gsm": _first(data, "gsm"),
|
||||
"ani": _first(data, "ani"),
|
||||
"uraProtocolId": _first(data, "uraProtocolId", "ura_protocol_id"),
|
||||
"uraLatency": _first(data, "uraLatency", "ura_latency"),
|
||||
"uraResolution": _first(data, "uraResolution", "urResolution", "ura_resolution"),
|
||||
"customerMessage": _first(data, "customerMessage", "customer_message"),
|
||||
# Message/guardrails/analysis
|
||||
"messageId": _first(data, "messageId", "message_id"),
|
||||
"blockingGuardrailsOutput": _first(data, "blockingGuardrailsOutput", "blocking_guardrails_output"),
|
||||
"blockingGuardrailsInput": _first(data, "blockingGuardrailsInput", "blocking_guardrails_input"),
|
||||
"llmResponse": _first(data, "llmResponse", "llm_response"),
|
||||
"alucinationScore": _first(data, "alucinationScore", "hallucinationScore", "alucination_score"),
|
||||
"noMatchRag": _first(data, "noMatchRag", "no_match_rag"),
|
||||
"promptLength": _first(data, "promptLength", "prompt_length"),
|
||||
"intention": _first(data, "intention", "intent"),
|
||||
"loop": _first(data, "loop"),
|
||||
"inferredCsiScore": _first(data, "inferredCsiScore", "inferred_csi_score"),
|
||||
"supervisorBlockReasons": _first(data, "supervisorBlockReasons", "supervisor_block_reasons"),
|
||||
"resolution": _first(data, "resolution"),
|
||||
"ConversationPrecision": _first(data, "ConversationPrecision", "conversationPrecision", "conversation_precision"),
|
||||
# LLM metrics
|
||||
"model": _first(data, "model") or event.get("model"),
|
||||
"tokenInput": _first(token_usage, "input_tokens") or _first(data, "tokenInput", "input_tokens"),
|
||||
"tokenOutput": _first(token_usage, "output_tokens") or _first(data, "tokenOutput", "output_tokens"),
|
||||
"latencyMs": _first(data, "latencyMs", "duration_ms"),
|
||||
"toxicityScore": _first(data, "toxicityScore", "toxicity_score"),
|
||||
"nps": _first(data, "nps"),
|
||||
"judgeScore": _first(data, "judgeScore", "judge_score"),
|
||||
"accuracyScore": _first(data, "accuracyScore", "accuracy_score"),
|
||||
"guardrails": _first(data, "guardrails"),
|
||||
# RAG
|
||||
"ragRetrievedDocuments": _as_list(_first(data, "documentsRetrieved", "ragRetrievedDocuments")),
|
||||
"ragSelectedDocuments": _as_list(_first(data, "documentsSelected", "ragSelectedDocuments")),
|
||||
# API
|
||||
"apiUrl": _first(data, "apiUrl", "api_url"),
|
||||
"apiStatusCode": _first(data, "httpStatusCode", "apiStatusCode", "http_status_code"),
|
||||
"apiResponsePayload": _first(data, "apiResponsePayload", "api_response_payload"),
|
||||
# I/O
|
||||
"inputData": _first(data, "inputData", "input_data"),
|
||||
"outputData": _first(data, "outputData", "output_data"),
|
||||
# Business/status/sequence
|
||||
"agentSpecificData": _collect_agent_specific_data(metadata, body),
|
||||
"status": _first(data, "status"),
|
||||
"sequence": _first(data, "sequence"),
|
||||
}
|
||||
|
||||
if keep_none:
|
||||
return {k: ("" if v is None else v) for k, v in payload.items()}
|
||||
return {k: v for k, v in payload.items() if v is not None}
|
||||
@@ -0,0 +1,396 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from collections import defaultdict
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Literal
|
||||
|
||||
logger = logging.getLogger("agent_framework.analytics.tim_sequence")
|
||||
|
||||
# In-process fallback. This is not cross-process/global, but keeps telemetry alive
|
||||
# when the configured shared sequence backend is unavailable, matching the
|
||||
# framework principle that observability must not break business execution.
|
||||
_memory_lock = threading.Lock()
|
||||
_memory_counters: dict[str, int] = defaultdict(int)
|
||||
|
||||
SequenceProvider = Literal["auto", "redis", "mongodb", "mongo", "memory", "none"]
|
||||
|
||||
|
||||
def _env_bool(name: str, default: bool) -> bool:
|
||||
value = os.getenv(name)
|
||||
if value is None:
|
||||
return default
|
||||
return value.strip().lower() in {"1", "true", "yes", "y", "on"}
|
||||
|
||||
|
||||
def sequence_enabled() -> bool:
|
||||
return _env_bool("PUBSUB_SEQUENCE_ENABLED", True)
|
||||
|
||||
|
||||
def _sequence_provider() -> SequenceProvider:
|
||||
raw = (os.getenv("PUBSUB_SEQUENCE_PROVIDER") or "auto").strip().lower()
|
||||
if raw in {"mongo"}:
|
||||
return "mongodb"
|
||||
if raw in {"auto", "redis", "mongodb", "memory", "none"}:
|
||||
return raw # type: ignore[return-value]
|
||||
logger.warning("tim_sequence.invalid_provider provider=%s; using auto", raw)
|
||||
return "auto"
|
||||
|
||||
|
||||
def _redis_url() -> str | None:
|
||||
return os.getenv("PUBSUB_SEQUENCE_REDIS_URL") or os.getenv("REDIS_URL")
|
||||
|
||||
|
||||
def _mongo_uri() -> str | None:
|
||||
return (
|
||||
os.getenv("PUBSUB_SEQUENCE_MONGODB_URI")
|
||||
or os.getenv("MONGODB_URI")
|
||||
or os.getenv("MONGO_URI")
|
||||
)
|
||||
|
||||
|
||||
def _mongo_database() -> str:
|
||||
return (
|
||||
os.getenv("PUBSUB_SEQUENCE_MONGODB_DATABASE")
|
||||
or os.getenv("MONGODB_DATABASE")
|
||||
or os.getenv("MONGO_DATABASE")
|
||||
or "agent_platform"
|
||||
)
|
||||
|
||||
|
||||
def _legacy_agent_name() -> str:
|
||||
return _safe_part(os.getenv("AGENT_NAME") or "agent", "agent")
|
||||
|
||||
|
||||
def _mongo_collection() -> str:
|
||||
"""Return the shared MongoDB collection used by every event producer.
|
||||
|
||||
The collection must not vary by agent. A transaction can emit GRL, AGA,
|
||||
NOC and other events from different components, and all of them must
|
||||
increment the same counter document. Deployments may override the name,
|
||||
but the configured value must be identical in every producer/pod.
|
||||
"""
|
||||
return (
|
||||
os.getenv("PUBSUB_SEQUENCE_MONGODB_COLLECTION")
|
||||
or os.getenv("MONGODB_EVENT_COUNTERS_COLLECTION")
|
||||
or os.getenv("EVENT_COUNTERS_COLLECTION")
|
||||
or "observer_event_counters"
|
||||
)
|
||||
|
||||
|
||||
def _ttl_seconds() -> int:
|
||||
raw = os.getenv("PUBSUB_SEQUENCE_TTL_SECONDS") or os.getenv("SESSION_TTL_SECONDS") or "86400"
|
||||
try:
|
||||
return max(0, int(raw))
|
||||
except Exception:
|
||||
return 86400
|
||||
|
||||
|
||||
def _fallback_enabled() -> bool:
|
||||
# An in-memory fallback creates duplicate sequences when multiple pods or
|
||||
# event producers handle the same transaction. Keep it opt-in only for
|
||||
# local/single-process development.
|
||||
return _env_bool("PUBSUB_SEQUENCE_MEMORY_FALLBACK", False)
|
||||
|
||||
|
||||
def _key_prefix() -> str:
|
||||
return os.getenv("PUBSUB_SEQUENCE_KEY_PREFIX") or "observer:sequence"
|
||||
|
||||
|
||||
def _safe_part(value: Any, fallback: str) -> str:
|
||||
text = str(value or fallback).strip()
|
||||
return text.replace(" ", "_").replace("/", "_").replace("\\", "_")
|
||||
|
||||
|
||||
def build_sequence_key(
|
||||
agent_id: str | None,
|
||||
session_id: str | None,
|
||||
transaction_id: str | None = None,
|
||||
) -> str:
|
||||
"""Build one counter key for the whole transaction.
|
||||
|
||||
``agent_id`` is intentionally ignored for transaction-scoped counters.
|
||||
A single transaction may emit events from different agents/components
|
||||
(for example GRL and AGA), and those events must share one monotonic
|
||||
sequence. ``session_id`` is retained only as a compatibility fallback when
|
||||
no transaction identifier is present.
|
||||
"""
|
||||
if transaction_id:
|
||||
transaction = _safe_part(transaction_id, "unknown_transaction")
|
||||
return f"{_key_prefix()}:transaction:{transaction}"
|
||||
|
||||
# Legacy fallback. Including the agent here avoids changing old session-only
|
||||
# behavior, but new integrations should always provide transactionId.
|
||||
agent = _safe_part(agent_id or os.getenv("AGENT_NAME"), "agent")
|
||||
session = _safe_part(session_id, "unknown_session")
|
||||
return f"{_key_prefix()}:{agent}:session:{session}"
|
||||
|
||||
|
||||
async def _next_sequence_redis(key: str, ttl_seconds: int) -> int | None:
|
||||
url = _redis_url()
|
||||
if not url:
|
||||
return None
|
||||
try:
|
||||
import redis.asyncio as redis_async # type: ignore
|
||||
|
||||
client = redis_async.Redis.from_url(url, decode_responses=True)
|
||||
try:
|
||||
value = await client.incr(key)
|
||||
if ttl_seconds > 0 and value == 1:
|
||||
await client.expire(key, ttl_seconds)
|
||||
return int(value)
|
||||
finally:
|
||||
try:
|
||||
await client.aclose()
|
||||
except AttributeError: # redis-py older compatibility
|
||||
await client.close()
|
||||
except Exception:
|
||||
logger.exception("tim_sequence.redis_failed key=%s", key)
|
||||
return None
|
||||
|
||||
|
||||
_mongo_index_checked = False
|
||||
_mongo_index_lock = threading.Lock()
|
||||
|
||||
|
||||
def _next_sequence_mongodb_sync(
|
||||
key: str,
|
||||
agent_id: str | None,
|
||||
session_id: str | None,
|
||||
transaction_id: str | None,
|
||||
ttl_seconds: int,
|
||||
) -> int | None:
|
||||
uri = _mongo_uri()
|
||||
if not uri:
|
||||
return None
|
||||
|
||||
from pymongo import MongoClient, ReturnDocument # type: ignore
|
||||
|
||||
client = MongoClient(uri)
|
||||
try:
|
||||
collection = client[_mongo_database()][_mongo_collection()]
|
||||
now = datetime.now(timezone.utc)
|
||||
expires_at = now + timedelta(seconds=ttl_seconds) if ttl_seconds > 0 else None
|
||||
|
||||
# update: dict[str, Any] = {
|
||||
# "$inc": {"sequence": 1},
|
||||
# "$set": {
|
||||
# "agentId": agent_id or os.getenv("AGENT_NAME") or "agent",
|
||||
# "sessionId": session_id,
|
||||
# "transactionId": transaction_id,
|
||||
# "sequenceScope": "transaction" if transaction_id else "session",
|
||||
# "updatedAt": now,
|
||||
# },
|
||||
# "$setOnInsert": {
|
||||
# "_id": key,
|
||||
# "createdAt": now,
|
||||
# },
|
||||
# }
|
||||
update: dict[str, Any] = {
|
||||
"$inc": {"sequence": 1},
|
||||
"$set": {
|
||||
"agentId": agent_id or os.getenv("AGENT_NAME") or "agent",
|
||||
"sessionId": session_id,
|
||||
"transactionId": transaction_id,
|
||||
"sequenceScope": "transaction" if transaction_id else "session",
|
||||
"updatedAt": now,
|
||||
},
|
||||
"$setOnInsert": {
|
||||
"createdAt": now,
|
||||
},
|
||||
}
|
||||
if expires_at is not None:
|
||||
update["$set"]["expiresAt"] = expires_at
|
||||
|
||||
doc = collection.find_one_and_update(
|
||||
{"_id": key},
|
||||
update,
|
||||
upsert=True,
|
||||
return_document=ReturnDocument.AFTER,
|
||||
)
|
||||
if not doc:
|
||||
return None
|
||||
return int(doc.get("sequence", 0))
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
|
||||
def _ensure_mongo_ttl_index_once_sync(ttl_seconds: int) -> None:
|
||||
"""Best-effort TTL index initialization, safe across threads/event loops.
|
||||
|
||||
``asyncio.Lock`` must not be shared by independent event loops. Observer
|
||||
compatibility calls may originate in worker threads, so this one-time
|
||||
process-local guard deliberately uses ``threading.Lock``. The blocking
|
||||
Mongo operation is executed by the async wrapper in a worker thread.
|
||||
"""
|
||||
global _mongo_index_checked
|
||||
if _mongo_index_checked or ttl_seconds <= 0 or not _mongo_uri():
|
||||
return
|
||||
|
||||
with _mongo_index_lock:
|
||||
if _mongo_index_checked:
|
||||
return
|
||||
try:
|
||||
from pymongo import MongoClient # type: ignore
|
||||
|
||||
client = MongoClient(_mongo_uri())
|
||||
try:
|
||||
collection = client[_mongo_database()][_mongo_collection()]
|
||||
collection.create_index("expiresAt", expireAfterSeconds=0, background=True)
|
||||
finally:
|
||||
client.close()
|
||||
except Exception:
|
||||
logger.warning("tim_sequence.mongodb_ttl_index_failed", exc_info=True)
|
||||
finally:
|
||||
# The index is an observability housekeeping concern, not a
|
||||
# prerequisite for sequence generation. Do not retry on every
|
||||
# event if the application user lacks index privileges.
|
||||
_mongo_index_checked = True
|
||||
|
||||
|
||||
async def _ensure_mongo_ttl_index_once(ttl_seconds: int) -> None:
|
||||
await asyncio.to_thread(_ensure_mongo_ttl_index_once_sync, ttl_seconds)
|
||||
|
||||
|
||||
async def _next_sequence_mongodb(
|
||||
key: str,
|
||||
agent_id: str | None,
|
||||
session_id: str | None,
|
||||
transaction_id: str | None,
|
||||
ttl_seconds: int,
|
||||
) -> int | None:
|
||||
if not _mongo_uri():
|
||||
return None
|
||||
try:
|
||||
await _ensure_mongo_ttl_index_once(ttl_seconds)
|
||||
return await asyncio.to_thread(
|
||||
_next_sequence_mongodb_sync,
|
||||
key,
|
||||
agent_id,
|
||||
session_id,
|
||||
transaction_id,
|
||||
ttl_seconds,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("tim_sequence.mongodb_failed key=%s", key)
|
||||
return None
|
||||
|
||||
|
||||
async def _next_sequence_memory(key: str) -> int:
|
||||
# Tiny in-process critical section; a thread lock is intentional because
|
||||
# this fallback can be reached from more than one asyncio event loop.
|
||||
with _memory_lock:
|
||||
_memory_counters[key] += 1
|
||||
return _memory_counters[key]
|
||||
|
||||
|
||||
async def next_sequence(
|
||||
agent_id: str | None,
|
||||
session_id: str | None,
|
||||
transaction_id: str | None = None,
|
||||
) -> int | None:
|
||||
"""Return the next observer sequence isolated by transaction.
|
||||
|
||||
The preferred scope is only ``transaction_id``. Agent/event family must
|
||||
never participate in the key because one transaction can emit events from
|
||||
several components. ``session_id`` is used only as a backward-compatible
|
||||
fallback. Redis and MongoDB increments remain atomic across replicas.
|
||||
"""
|
||||
if not sequence_enabled() or (not transaction_id and not session_id):
|
||||
return None
|
||||
|
||||
provider = _sequence_provider()
|
||||
if provider == "none":
|
||||
return None
|
||||
|
||||
key = build_sequence_key(agent_id, session_id, transaction_id)
|
||||
ttl_seconds = _ttl_seconds()
|
||||
value: int | None = None
|
||||
|
||||
if provider == "memory":
|
||||
return await _next_sequence_memory(key)
|
||||
|
||||
if provider == "redis":
|
||||
value = await _next_sequence_redis(key, ttl_seconds)
|
||||
elif provider == "mongodb":
|
||||
value = await _next_sequence_mongodb(
|
||||
key, agent_id, session_id, transaction_id, ttl_seconds
|
||||
)
|
||||
else: # auto
|
||||
if _redis_url():
|
||||
value = await _next_sequence_redis(key, ttl_seconds)
|
||||
if value is None and _mongo_uri():
|
||||
value = await _next_sequence_mongodb(
|
||||
key, agent_id, session_id, transaction_id, ttl_seconds
|
||||
)
|
||||
|
||||
if value is not None:
|
||||
return value
|
||||
if _fallback_enabled():
|
||||
return await _next_sequence_memory(key)
|
||||
return None
|
||||
|
||||
|
||||
async def ensure_sequence(payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Inject sequence if missing, preserving explicit values from metadata/body.
|
||||
|
||||
Used by the flat Pub/Sub schema, where sessionId/agentId sit at the root.
|
||||
For the nested analytics envelope (OCI Streaming) use
|
||||
:func:`ensure_sequence_envelope`.
|
||||
"""
|
||||
if not isinstance(payload, dict):
|
||||
return payload
|
||||
if payload.get("sequence") is not None:
|
||||
return payload
|
||||
session_id = payload.get("sessionId") or payload.get("session_id")
|
||||
transaction_id = (
|
||||
payload.get("transactionId")
|
||||
or payload.get("transaction_id")
|
||||
or payload.get("transactionID")
|
||||
)
|
||||
agent_id = payload.get("agentId") or payload.get("agent_id") or os.getenv("AGENT_NAME")
|
||||
seq = await next_sequence(agent_id, session_id, transaction_id)
|
||||
if seq is not None:
|
||||
payload["sequence"] = seq
|
||||
return payload
|
||||
|
||||
|
||||
async def ensure_sequence_envelope(event: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Inject sequence into a ``build_analytics_event`` envelope.
|
||||
|
||||
The envelope shape is ``{eventType, source, eventDate, payload, metadata}``.
|
||||
Unlike the flat Pub/Sub payload, sessionId/agentId are not at the root: they
|
||||
live inside ``payload`` and/or ``metadata``. We read them from the merged
|
||||
``{**payload, **metadata}`` view, mirroring the flat mapper
|
||||
(tim_payload_mapper.map_analytics_event_to_tim_flat_payload) and the legacy
|
||||
observer (observer/api.py: metadata.sessionId -> sessionId).
|
||||
|
||||
The counter is written at the envelope root, as a sibling of ``eventType`` —
|
||||
the faithful analog of the legacy flat payload where ``sequence`` sat next to
|
||||
``eventType``/``traceId``. The outer transport contract ``{type, payload}`` is
|
||||
left untouched; only this inner field is added.
|
||||
"""
|
||||
if not isinstance(event, dict):
|
||||
return event
|
||||
if event.get("sequence") is not None:
|
||||
return event
|
||||
body = event.get("payload") if isinstance(event.get("payload"), dict) else {}
|
||||
metadata = event.get("metadata") if isinstance(event.get("metadata"), dict) else {}
|
||||
data = {**body, **metadata}
|
||||
session_id = data.get("sessionId") or data.get("session_id")
|
||||
# Os adapters do BO emitem snake_case; o contrato TIM usa transactionId e
|
||||
# payloads antigos trazem transactionID. Sem as tres grafias o contador cai
|
||||
# em escopo de sessao e perde o isolamento por transacao.
|
||||
transaction_id = (
|
||||
data.get("transactionId")
|
||||
or data.get("transaction_id")
|
||||
or data.get("transactionID")
|
||||
)
|
||||
agent_id = data.get("agentId") or data.get("agent_id") or os.getenv("AGENT_NAME")
|
||||
seq = await next_sequence(agent_id, session_id, transaction_id)
|
||||
if seq is not None:
|
||||
event["sequence"] = seq
|
||||
return event
|
||||
@@ -0,0 +1 @@
|
||||
from .usage_repository import UsageRecord, UsageRepository, SQLiteUsageRepository, OracleUsageRepository, create_usage_repository
|
||||
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,173 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from dataclasses import dataclass, asdict
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from agent_framework.observability.context import get_observability_context
|
||||
|
||||
@dataclass
|
||||
class UsageRecord:
|
||||
provider: str
|
||||
model: str
|
||||
operation: str
|
||||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
cached_tokens: int = 0
|
||||
total_tokens: int = 0
|
||||
cost_usd: float = 0.0
|
||||
cost_brl: float = 0.0
|
||||
metadata: dict[str, Any] | None = None
|
||||
request_id: str | None = None
|
||||
session_id: str | None = None
|
||||
tenant_id: str | None = None
|
||||
agent_id: str | None = None
|
||||
user_id: str | None = None
|
||||
message_id: str | None = None
|
||||
created_at: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_usage(cls, provider: str, model: str, operation: str, usage: dict[str, Any], metadata: dict[str, Any] | None = None) -> "UsageRecord":
|
||||
ctx = get_observability_context()
|
||||
return cls(
|
||||
provider=provider, model=model, operation=operation,
|
||||
prompt_tokens=int(usage.get("prompt_tokens") or 0),
|
||||
completion_tokens=int(usage.get("completion_tokens") or 0),
|
||||
cached_tokens=int(usage.get("cached_tokens") or 0),
|
||||
total_tokens=int(usage.get("total_tokens") or 0),
|
||||
cost_usd=float(usage.get("cost_usd") or 0),
|
||||
cost_brl=float(usage.get("cost_brl") or 0),
|
||||
metadata=metadata or {}, request_id=ctx.request_id, session_id=ctx.session_id,
|
||||
tenant_id=ctx.tenant_id, agent_id=ctx.agent_id, user_id=ctx.user_id,
|
||||
message_id=ctx.message_id, created_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
def model_dump(self) -> dict[str, Any]:
|
||||
return asdict(self)
|
||||
|
||||
class UsageRepository:
|
||||
async def record(self, usage: UsageRecord) -> None: ...
|
||||
async def summarize(self, *, tenant_id: str | None = None, session_id: str | None = None) -> dict[str, Any]: ...
|
||||
|
||||
class SQLiteUsageRepository(UsageRepository):
|
||||
def __init__(self, settings):
|
||||
from agent_framework.persistence.sqlite_store import SQLiteStore
|
||||
self.store = SQLiteStore(settings.SQLITE_DB_PATH)
|
||||
self._init_schema()
|
||||
|
||||
def _init_schema(self):
|
||||
ddl = """
|
||||
create table if not exists llm_usage_records (
|
||||
id integer primary key autoincrement,
|
||||
request_id text, session_id text, tenant_id text, agent_id text, user_id text, message_id text,
|
||||
provider text not null, model text not null, operation text not null,
|
||||
prompt_tokens integer not null default 0,
|
||||
completion_tokens integer not null default 0,
|
||||
cached_tokens integer not null default 0,
|
||||
total_tokens integer not null default 0,
|
||||
cost_usd real not null default 0,
|
||||
cost_brl real not null default 0,
|
||||
metadata_json text,
|
||||
created_at text not null
|
||||
);
|
||||
create index if not exists idx_usage_tenant_created on llm_usage_records(tenant_id, created_at);
|
||||
create index if not exists idx_usage_session_created on llm_usage_records(session_id, created_at);
|
||||
"""
|
||||
with self.store._lock, self.store.connect() as con:
|
||||
con.executescript(ddl)
|
||||
|
||||
async def record(self, usage: UsageRecord) -> None:
|
||||
with self.store._lock, self.store.connect() as con:
|
||||
con.execute("""
|
||||
insert into llm_usage_records(
|
||||
request_id,session_id,tenant_id,agent_id,user_id,message_id,
|
||||
provider,model,operation,prompt_tokens,completion_tokens,cached_tokens,total_tokens,
|
||||
cost_usd,cost_brl,metadata_json,created_at
|
||||
) values(?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)
|
||||
""", (
|
||||
usage.request_id, usage.session_id, usage.tenant_id, usage.agent_id, usage.user_id, usage.message_id,
|
||||
usage.provider, usage.model, usage.operation, usage.prompt_tokens, usage.completion_tokens,
|
||||
usage.cached_tokens, usage.total_tokens, usage.cost_usd, usage.cost_brl,
|
||||
json.dumps(usage.metadata or {}, ensure_ascii=False, default=str), usage.created_at,
|
||||
))
|
||||
|
||||
async def summarize(self, *, tenant_id: str | None = None, session_id: str | None = None) -> dict[str, Any]:
|
||||
where=[]; params=[]
|
||||
if tenant_id: where.append('tenant_id=?'); params.append(tenant_id)
|
||||
if session_id: where.append('session_id=?'); params.append(session_id)
|
||||
sql="""select count(*) calls, coalesce(sum(prompt_tokens),0) prompt_tokens,
|
||||
coalesce(sum(completion_tokens),0) completion_tokens,
|
||||
coalesce(sum(total_tokens),0) total_tokens,
|
||||
coalesce(sum(cost_usd),0) cost_usd,
|
||||
coalesce(sum(cost_brl),0) cost_brl
|
||||
from llm_usage_records"""
|
||||
if where: sql += ' where ' + ' and '.join(where)
|
||||
with self.store._lock, self.store.connect() as con:
|
||||
row=con.execute(sql, params).fetchone()
|
||||
return dict(row) if row else {"calls":0,"prompt_tokens":0,"completion_tokens":0,"total_tokens":0,"cost_usd":0,"cost_brl":0}
|
||||
|
||||
class OracleUsageRepository(UsageRepository):
|
||||
def __init__(self, settings):
|
||||
from agent_framework.persistence.oracle_store import OracleStore
|
||||
self.store = OracleStore(settings)
|
||||
self._init_schema()
|
||||
|
||||
def _init_schema(self):
|
||||
with self.store.connect() as conn:
|
||||
cur=conn.cursor()
|
||||
self.store._exec_ddl_ignore_exists(cur, f"""
|
||||
create table {self.store.t('LLM_USAGE_RECORD')} (
|
||||
ID number generated always as identity primary key,
|
||||
REQUEST_ID varchar2(128), SESSION_ID varchar2(256), TENANT_ID varchar2(128),
|
||||
AGENT_ID varchar2(128), USER_ID varchar2(256), MESSAGE_ID varchar2(256),
|
||||
PROVIDER varchar2(128) not null, MODEL varchar2(256) not null, OPERATION varchar2(128) not null,
|
||||
PROMPT_TOKENS number default 0, COMPLETION_TOKENS number default 0, CACHED_TOKENS number default 0,
|
||||
TOTAL_TOKENS number default 0, COST_USD number default 0, COST_BRL number default 0,
|
||||
METADATA_JSON clob check (METADATA_JSON is json), CREATED_AT timestamp with time zone not null
|
||||
)
|
||||
""")
|
||||
self.store._exec_ddl_ignore_exists(cur, f"create index {self.store.t('IX_USAGE_TENANT')} on {self.store.t('LLM_USAGE_RECORD')}(TENANT_ID, CREATED_AT)")
|
||||
self.store._exec_ddl_ignore_exists(cur, f"create index {self.store.t('IX_USAGE_SESSION')} on {self.store.t('LLM_USAGE_RECORD')}(SESSION_ID, CREATED_AT)")
|
||||
|
||||
async def record(self, usage: UsageRecord) -> None:
|
||||
await asyncio.to_thread(self._record_sync, usage)
|
||||
|
||||
def _record_sync(self, usage: UsageRecord):
|
||||
with self.store.connect() as conn:
|
||||
conn.cursor().execute(f"""
|
||||
insert into {self.store.t('LLM_USAGE_RECORD')}(
|
||||
REQUEST_ID,SESSION_ID,TENANT_ID,AGENT_ID,USER_ID,MESSAGE_ID,PROVIDER,MODEL,OPERATION,
|
||||
PROMPT_TOKENS,COMPLETION_TOKENS,CACHED_TOKENS,TOTAL_TOKENS,COST_USD,COST_BRL,METADATA_JSON,CREATED_AT
|
||||
) values(:1,:2,:3,:4,:5,:6,:7,:8,:9,:10,:11,:12,:13,:14,:15,:16,:17)
|
||||
""", [
|
||||
usage.request_id, usage.session_id, usage.tenant_id, usage.agent_id, usage.user_id, usage.message_id,
|
||||
usage.provider, usage.model, usage.operation, usage.prompt_tokens, usage.completion_tokens, usage.cached_tokens,
|
||||
usage.total_tokens, usage.cost_usd, usage.cost_brl, json.dumps(usage.metadata or {}, ensure_ascii=False, default=str), usage.created_at,
|
||||
])
|
||||
|
||||
async def summarize(self, *, tenant_id: str | None = None, session_id: str | None = None) -> dict[str, Any]:
|
||||
return await asyncio.to_thread(self._summarize_sync, tenant_id, session_id)
|
||||
|
||||
def _summarize_sync(self, tenant_id, session_id):
|
||||
where=[]; params={}
|
||||
if tenant_id: where.append('TENANT_ID=:tenant_id'); params['tenant_id']=tenant_id
|
||||
if session_id: where.append('SESSION_ID=:session_id'); params['session_id']=session_id
|
||||
sql=f"""select count(*) CALLS, coalesce(sum(PROMPT_TOKENS),0) PROMPT_TOKENS,
|
||||
coalesce(sum(COMPLETION_TOKENS),0) COMPLETION_TOKENS,
|
||||
coalesce(sum(TOTAL_TOKENS),0) TOTAL_TOKENS,
|
||||
coalesce(sum(COST_USD),0) COST_USD,
|
||||
coalesce(sum(COST_BRL),0) COST_BRL
|
||||
from {self.store.t('LLM_USAGE_RECORD')}"""
|
||||
if where: sql += ' where ' + ' and '.join(where)
|
||||
with self.store.connect() as conn:
|
||||
cur=conn.cursor(); cur.execute(sql, params); row=cur.fetchone()
|
||||
cols=[d[0].lower() for d in cur.description]
|
||||
return dict(zip(cols,row)) if row else {}
|
||||
|
||||
def create_usage_repository(settings) -> UsageRepository:
|
||||
provider = getattr(settings, 'USAGE_REPOSITORY_PROVIDER', None) or getattr(settings, 'MEMORY_REPOSITORY_PROVIDER', 'memory')
|
||||
if provider in {'autonomous','oracle'}:
|
||||
return OracleUsageRepository(settings)
|
||||
return SQLiteUsageRepository(settings)
|
||||
0
agent_framework_oci/libs/agent_framework/build/lib/agent_framework/cache/__init__.py
vendored
Normal file
0
agent_framework_oci/libs/agent_framework/build/lib/agent_framework/cache/__init__.py
vendored
Normal file
Binary file not shown.
Binary file not shown.
184
agent_framework_oci/libs/agent_framework/build/lib/agent_framework/cache/cache.py
vendored
Normal file
184
agent_framework_oci/libs/agent_framework/build/lib/agent_framework/cache/cache.py
vendored
Normal file
@@ -0,0 +1,184 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from datetime import datetime, timezone, timedelta
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger("agent_framework.cache")
|
||||
|
||||
|
||||
class Cache:
|
||||
async def get(self, key: str) -> Any | None: ...
|
||||
async def set(self, key: str, value: Any, ttl_seconds: int | None = None) -> None: ...
|
||||
async def delete(self, key: str) -> None: ...
|
||||
|
||||
|
||||
class InMemoryCache(Cache):
|
||||
def __init__(self):
|
||||
self._data: dict[str, tuple[Any, float | None]] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
async def get(self, key):
|
||||
async with self._lock:
|
||||
item = self._data.get(key)
|
||||
if not item:
|
||||
return None
|
||||
value, expires = item
|
||||
if expires and expires < time.time():
|
||||
self._data.pop(key, None)
|
||||
return None
|
||||
return value
|
||||
|
||||
async def set(self, key, value, ttl_seconds=None):
|
||||
async with self._lock:
|
||||
self._data[key] = (value, time.time() + ttl_seconds if ttl_seconds else None)
|
||||
|
||||
async def delete(self, key):
|
||||
async with self._lock:
|
||||
self._data.pop(key, None)
|
||||
|
||||
|
||||
class RedisCache(Cache):
|
||||
"""Redis L2 cache with redis-py sync/async compatibility and safe fallback."""
|
||||
def __init__(self, settings):
|
||||
self.url = settings.REDIS_URL
|
||||
self.prefix = getattr(settings, "CACHE_KEY_PREFIX", "agentfw")
|
||||
self._async = False
|
||||
try:
|
||||
import redis.asyncio as redis_async
|
||||
self.client = redis_async.Redis.from_url(self.url, decode_responses=True)
|
||||
self._async = True
|
||||
except Exception:
|
||||
import redis
|
||||
self.client = redis.Redis.from_url(self.url, decode_responses=True)
|
||||
|
||||
def _key(self, key: str) -> str:
|
||||
return f"{self.prefix}:{key}"
|
||||
|
||||
async def get(self, key):
|
||||
try:
|
||||
raw = await self.client.get(self._key(key)) if self._async else await asyncio.to_thread(self.client.get, self._key(key))
|
||||
return json.loads(raw) if raw else None
|
||||
except Exception:
|
||||
logger.exception("Redis GET falhou key=%s", key)
|
||||
return None
|
||||
|
||||
async def set(self, key, value, ttl_seconds=None):
|
||||
raw = json.dumps(value, ensure_ascii=False, default=str)
|
||||
try:
|
||||
if self._async:
|
||||
await self.client.set(self._key(key), raw, ex=ttl_seconds)
|
||||
else:
|
||||
await asyncio.to_thread(self.client.set, self._key(key), raw, ex=ttl_seconds)
|
||||
except Exception:
|
||||
logger.exception("Redis SET falhou key=%s", key)
|
||||
|
||||
async def delete(self, key):
|
||||
try:
|
||||
if self._async:
|
||||
await self.client.delete(self._key(key))
|
||||
else:
|
||||
await asyncio.to_thread(self.client.delete, self._key(key))
|
||||
except Exception:
|
||||
logger.exception("Redis DELETE falhou key=%s", key)
|
||||
|
||||
|
||||
class SQLiteCache(Cache):
|
||||
def __init__(self, settings):
|
||||
from agent_framework.persistence.sqlite_store import SQLiteStore
|
||||
self.store = SQLiteStore(settings.SQLITE_DB_PATH)
|
||||
|
||||
async def get(self, key):
|
||||
return await asyncio.to_thread(self._get_sync, key)
|
||||
|
||||
def _get_sync(self, key):
|
||||
with self.store._lock, self.store.connect() as con:
|
||||
row = con.execute("select value_json, expires_at from cache_entries where key=?", (key,)).fetchone()
|
||||
if not row:
|
||||
return None
|
||||
if row["expires_at"] and row["expires_at"] < time.time():
|
||||
con.execute("delete from cache_entries where key=?", (key,))
|
||||
return None
|
||||
return json.loads(row["value_json"])
|
||||
|
||||
async def set(self, key, value, ttl_seconds=None):
|
||||
await asyncio.to_thread(self._set_sync, key, value, ttl_seconds)
|
||||
|
||||
def _set_sync(self, key, value, ttl_seconds=None):
|
||||
expires = time.time() + ttl_seconds if ttl_seconds else None
|
||||
with self.store._lock, self.store.connect() as con:
|
||||
con.execute(
|
||||
"insert or replace into cache_entries(key,value_json,expires_at,created_at) values(?,?,?,?)",
|
||||
(key, json.dumps(value, ensure_ascii=False, default=str), expires, self.store.now()),
|
||||
)
|
||||
|
||||
async def delete(self, key):
|
||||
await asyncio.to_thread(self._delete_sync, key)
|
||||
|
||||
def _delete_sync(self, key):
|
||||
with self.store._lock, self.store.connect() as con:
|
||||
con.execute("delete from cache_entries where key=?", (key,))
|
||||
|
||||
|
||||
class OracleCache(Cache):
|
||||
def __init__(self, settings):
|
||||
from agent_framework.persistence.oracle_store import OracleStore
|
||||
self.store = OracleStore(settings)
|
||||
|
||||
async def get(self, key): return await self.store.cache_get(key)
|
||||
async def set(self, key, value, ttl_seconds=None):
|
||||
expires = datetime.now(timezone.utc) + timedelta(seconds=ttl_seconds) if ttl_seconds else None
|
||||
await self.store.cache_set(key, value, expires_at=expires)
|
||||
async def delete(self, key): await self.store.cache_delete(key)
|
||||
|
||||
|
||||
class DistributedCache(Cache):
|
||||
"""L1 memory + optional L2 Redis/SQLite/Oracle with telemetry hooks."""
|
||||
def __init__(self, l1: Cache, l2: Cache | None = None, telemetry=None, default_ttl: int | None = None):
|
||||
self.l1, self.l2, self.telemetry, self.default_ttl = l1, l2, telemetry, default_ttl
|
||||
|
||||
async def get(self, key):
|
||||
v = await self.l1.get(key)
|
||||
if v is not None:
|
||||
if self.telemetry: await self.telemetry.cache_event("hit.l1", key, True)
|
||||
return v
|
||||
if not self.l2:
|
||||
if self.telemetry: await self.telemetry.cache_event("miss", key, False)
|
||||
return None
|
||||
v = await self.l2.get(key)
|
||||
if v is not None:
|
||||
await self.l1.set(key, v, self.default_ttl)
|
||||
if self.telemetry: await self.telemetry.cache_event("hit.l2", key, True)
|
||||
return v
|
||||
if self.telemetry: await self.telemetry.cache_event("miss", key, False)
|
||||
return None
|
||||
|
||||
async def set(self, key, value, ttl_seconds=None):
|
||||
ttl = ttl_seconds if ttl_seconds is not None else self.default_ttl
|
||||
await self.l1.set(key, value, ttl)
|
||||
if self.l2: await self.l2.set(key, value, ttl)
|
||||
if self.telemetry: await self.telemetry.cache_event("set", key, None, {"ttl_seconds": ttl})
|
||||
|
||||
async def delete(self, key):
|
||||
await self.l1.delete(key)
|
||||
if self.l2: await self.l2.delete(key)
|
||||
if self.telemetry: await self.telemetry.cache_event("delete", key, None)
|
||||
|
||||
|
||||
def create_cache(settings, telemetry=None):
|
||||
l1 = InMemoryCache()
|
||||
l2 = None
|
||||
if getattr(settings, "ENABLE_REDIS_CACHE", False):
|
||||
try:
|
||||
l2 = RedisCache(settings)
|
||||
except Exception:
|
||||
logger.exception("Redis indisponível; cache seguirá apenas com L1 memória")
|
||||
l2 = None
|
||||
if l2 is None:
|
||||
provider = getattr(settings, "CACHE_BACKEND_PROVIDER", "memory")
|
||||
if provider == "sqlite": l2 = SQLiteCache(settings)
|
||||
elif provider in {"autonomous", "oracle"}: l2 = OracleCache(settings)
|
||||
return DistributedCache(l1, l2, telemetry=telemetry, default_ttl=getattr(settings, "CACHE_TTL_SECONDS", None))
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,69 @@
|
||||
from .base import ChannelAdapter, ChannelMessage, ChannelResponse
|
||||
|
||||
|
||||
def _merge_context(payload: dict) -> dict:
|
||||
"""Preserva todo payload como contexto.
|
||||
|
||||
Antes o WebAdapter só copiava payload["context"]. Com isso, campos como
|
||||
business_context, msisdn, invoice_id e ura_call_id eram perdidos antes de
|
||||
chegar ao workflow/MCP.
|
||||
"""
|
||||
payload = dict(payload or {})
|
||||
ctx = dict(payload.get("context") or {})
|
||||
for k, v in payload.items():
|
||||
if k != "context" and k not in ctx:
|
||||
ctx[k] = v
|
||||
return ctx
|
||||
|
||||
|
||||
class WebAdapter(ChannelAdapter):
|
||||
name = "web"
|
||||
|
||||
async def normalize(self, payload):
|
||||
payload = payload or {}
|
||||
text = payload.get("message") or payload.get("text") or payload.get("content") or ""
|
||||
return ChannelMessage(
|
||||
channel="web",
|
||||
text=text,
|
||||
session_id=payload.get("session_id"),
|
||||
user_id=payload.get("user_id"),
|
||||
channel_id=payload.get("channel_id") or payload.get("channelId"),
|
||||
context=_merge_context(payload),
|
||||
)
|
||||
|
||||
async def render(self, response):
|
||||
return response.model_dump()
|
||||
|
||||
|
||||
class WhatsAppAdapter(ChannelAdapter):
|
||||
name = "whatsapp"
|
||||
|
||||
async def normalize(self, payload):
|
||||
payload = payload or {}
|
||||
return ChannelMessage(
|
||||
channel="whatsapp",
|
||||
channel_id=payload.get("from"),
|
||||
text=payload.get("text") or payload.get("message") or "",
|
||||
session_id=payload.get("session_id"),
|
||||
context=_merge_context(payload),
|
||||
)
|
||||
|
||||
async def render(self, response):
|
||||
return {"to": response.metadata.get("channel_id"), "text": response.text, "session_id": response.session_id}
|
||||
|
||||
|
||||
class VoiceAdapter(ChannelAdapter):
|
||||
name = "voice"
|
||||
|
||||
async def normalize(self, payload):
|
||||
payload = payload or {}
|
||||
return ChannelMessage(
|
||||
channel="voice",
|
||||
channel_id=payload.get("ani"),
|
||||
text=payload.get("transcript") or payload.get("text") or payload.get("message") or "",
|
||||
session_id=payload.get("session_id"),
|
||||
context=_merge_context(payload),
|
||||
)
|
||||
|
||||
async def render(self, response):
|
||||
return {"speak": response.text, "session_id": response.session_id}
|
||||
@@ -0,0 +1,21 @@
|
||||
from pydantic import BaseModel, Field
|
||||
from typing import Any
|
||||
|
||||
class ChannelMessage(BaseModel):
|
||||
channel: str
|
||||
channel_id: str | None = None
|
||||
session_id: str | None = None
|
||||
user_id: str | None = None
|
||||
text: str
|
||||
context: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
class ChannelResponse(BaseModel):
|
||||
channel: str
|
||||
session_id: str
|
||||
text: str
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
class ChannelAdapter:
|
||||
name = 'base'
|
||||
async def normalize(self, payload: dict) -> ChannelMessage: ...
|
||||
async def render(self, response: ChannelResponse) -> dict: ...
|
||||
@@ -0,0 +1,92 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from .adapters import WebAdapter, WhatsAppAdapter, VoiceAdapter, _merge_context
|
||||
from .base import ChannelMessage, ChannelResponse
|
||||
|
||||
try:
|
||||
from agent_framework.config.settings import settings
|
||||
except Exception: # pragma: no cover
|
||||
settings = None
|
||||
|
||||
|
||||
class ChannelGateway:
|
||||
"""Normalize and render messages at the Agent Framework boundary.
|
||||
|
||||
This class is used by the Agent Framework backend, not by the external
|
||||
Channel Gateway service.
|
||||
|
||||
input_mode semantics:
|
||||
- embedded: the backend may use internal channel adapters to interpret
|
||||
simple/native channel payloads. This is useful for demos, labs and local
|
||||
testing.
|
||||
- external: the backend expects a GatewayRequest payload that was already
|
||||
normalized by an external Channel Gateway. In this mode the backend does
|
||||
not parse native WhatsApp, Voice, Teams, or other channel payloads.
|
||||
|
||||
Backward compatibility:
|
||||
- The legacy constructor argument ``mode`` and setting
|
||||
``CHANNEL_GATEWAY_MODE`` are still accepted, but the preferred setting is
|
||||
``FRAMEWORK_CHANNEL_INPUT_MODE``.
|
||||
"""
|
||||
|
||||
def __init__(self, input_mode: str | None = None, mode: str | None = None):
|
||||
configured = (
|
||||
input_mode
|
||||
or mode
|
||||
or getattr(settings, "FRAMEWORK_CHANNEL_INPUT_MODE", None)
|
||||
or getattr(settings, "CHANNEL_GATEWAY_MODE", None)
|
||||
or "embedded"
|
||||
)
|
||||
self.input_mode = str(configured).strip().lower()
|
||||
if self.input_mode not in {"embedded", "external"}:
|
||||
raise ValueError(
|
||||
"INVALID_FRAMEWORK_CHANNEL_INPUT_MODE: expected 'embedded' or 'external'"
|
||||
)
|
||||
# Compatibility with previous code that accessed gateway.mode.
|
||||
self.mode = self.input_mode
|
||||
self.adapters = {a.name: a for a in [WebAdapter(), WhatsAppAdapter(), VoiceAdapter()]}
|
||||
|
||||
def get(self, channel: str):
|
||||
return self.adapters.get(channel, self.adapters["web"])
|
||||
|
||||
def _validate_external_payload(self, channel: str, payload: dict):
|
||||
"""Validate the payload portion of a GatewayRequest.
|
||||
|
||||
In external input mode, the backend is not accepting native channel
|
||||
payloads. It expects req.channel plus req.payload.message at minimum.
|
||||
Business keys remain optional because some journeys start without all
|
||||
identifiers and are completed by IdentityResolver or the agent.
|
||||
"""
|
||||
if not isinstance(channel, str) or not channel.strip():
|
||||
raise ValueError("INVALID_GATEWAY_REQUEST: channel is required")
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("INVALID_GATEWAY_REQUEST: payload must be an object")
|
||||
message = payload.get("message")
|
||||
if not isinstance(message, str) or not message.strip():
|
||||
raise ValueError(
|
||||
"INVALID_GATEWAY_REQUEST: payload.message is required and must be a non-empty string"
|
||||
)
|
||||
|
||||
async def _normalize_external(self, channel: str, payload: dict) -> ChannelMessage:
|
||||
self._validate_external_payload(channel, payload)
|
||||
return ChannelMessage(
|
||||
channel=channel,
|
||||
text=payload.get("message"),
|
||||
session_id=payload.get("session_id") or payload.get("session_key"),
|
||||
user_id=payload.get("user_id"),
|
||||
channel_id=payload.get("channel_id") or payload.get("channelId"),
|
||||
context=_merge_context(payload),
|
||||
)
|
||||
|
||||
async def normalize(self, channel: str, payload: dict) -> ChannelMessage:
|
||||
if self.input_mode == "external":
|
||||
return await self._normalize_external(channel, payload)
|
||||
return await self.get(channel).normalize(payload)
|
||||
|
||||
async def render(self, response: ChannelResponse) -> dict:
|
||||
if self.input_mode == "external":
|
||||
# The external Channel Gateway owns the final translation back to
|
||||
# WhatsApp, Voice, Teams, etc. The backend returns its canonical
|
||||
# response shape.
|
||||
return response.model_dump()
|
||||
return await self.get(response.channel).render(response)
|
||||
@@ -0,0 +1,156 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class InterruptionDecision:
|
||||
action: str # process | replay | classify
|
||||
text: str
|
||||
replay_text: str = ""
|
||||
reason: str = ""
|
||||
is_interruptible: bool = True
|
||||
terminal_status: str = ""
|
||||
heard_text: str = ""
|
||||
|
||||
|
||||
def _idle_nudges(payload: dict[str, Any]) -> list[str]:
|
||||
out: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for event in payload.get("events") or []:
|
||||
if not isinstance(event, dict) or event.get("type") != "idle_nudge":
|
||||
continue
|
||||
text = str(event.get("text") or "").strip()
|
||||
if text and text not in seen:
|
||||
seen.add(text)
|
||||
out.append(text)
|
||||
return out
|
||||
|
||||
|
||||
async def classify_processing_interruption(
|
||||
llm: Any,
|
||||
*,
|
||||
original_agent: str,
|
||||
original_client: str = "",
|
||||
supplement_client: str = "",
|
||||
profile_name: str = "processing_interruption_classifier",
|
||||
) -> bool:
|
||||
"""Decide se um barge-in interrompível exige regeneração da resposta.
|
||||
|
||||
Fail-safe: qualquer erro, resposta vazia ou formato inesperado retorna False,
|
||||
fazendo replay da fala anterior. O domínio não conhece este classificador;
|
||||
ele usa exclusivamente o LLMProvider do framework.
|
||||
"""
|
||||
if llm is None:
|
||||
return False
|
||||
prompt = (
|
||||
"Você classifica interrupções de voz durante uma resposta de atendimento. "
|
||||
"Responda somente 1 ou 0.\n"
|
||||
"1 = a fala/complemento do cliente adiciona ou altera informação relevante e "
|
||||
"a resposta do agente deve ser regenerada.\n"
|
||||
"0 = a interrupção não exige nova resposta; a fala anterior deve ser repetida.\n\n"
|
||||
f"Última fala do agente: {original_agent}\n"
|
||||
f"Última fala do cliente antes da resposta: {original_client}\n"
|
||||
f"Complemento/interrupção atual: {supplement_client}\n"
|
||||
)
|
||||
try:
|
||||
response = await llm.ainvoke(
|
||||
[{"role": "system", "content": prompt}],
|
||||
temperature=0,
|
||||
max_tokens=8,
|
||||
profile_name=profile_name,
|
||||
component_name=profile_name,
|
||||
generation_name=f"llm.{profile_name}",
|
||||
)
|
||||
raw = getattr(response, "content", response)
|
||||
text = str(raw or "").strip()
|
||||
return text.startswith("1")
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def evaluate_interruption(
|
||||
*,
|
||||
payload: dict[str, Any],
|
||||
message_text: str,
|
||||
session_metadata: dict[str, Any] | None,
|
||||
terminal_fallback_text: str = "",
|
||||
terminal_fallback_status: str = "erro_falha_sistema",
|
||||
) -> InterruptionDecision:
|
||||
"""Framework-level replay/interruption policy.
|
||||
|
||||
- sessão terminal: replay da última fala/fallback, sem reabrir o workflow;
|
||||
- idle_nudge: replay da última fala real;
|
||||
- fala não interrompível: replay;
|
||||
- fala interrompível com fala anterior: classificar antes de regenerar;
|
||||
- sem contexto anterior suficiente: processar normalmente.
|
||||
"""
|
||||
metadata = session_metadata or {}
|
||||
last_text = str(metadata.get("last_assistant_text") or "").strip()
|
||||
last_interruptible = bool(metadata.get("last_assistant_is_interruptible", True))
|
||||
|
||||
if bool(metadata.get("conversation_closed")):
|
||||
replay_text = (
|
||||
last_text
|
||||
or str(metadata.get("terminal_replay_text") or "").strip()
|
||||
or str(terminal_fallback_text or "").strip()
|
||||
)
|
||||
terminal_status = str(metadata.get("terminal_status") or "").strip() or terminal_fallback_status
|
||||
if replay_text:
|
||||
return InterruptionDecision(
|
||||
action="replay",
|
||||
text=message_text,
|
||||
replay_text=replay_text,
|
||||
reason="post_finalize",
|
||||
is_interruptible=False,
|
||||
terminal_status=terminal_status,
|
||||
)
|
||||
|
||||
if _idle_nudges(payload) and last_text:
|
||||
return InterruptionDecision(
|
||||
action="replay",
|
||||
text=message_text,
|
||||
replay_text=last_text,
|
||||
reason="idle_nudge",
|
||||
is_interruptible=last_interruptible,
|
||||
)
|
||||
|
||||
interruption = payload.get("processing_interruption")
|
||||
if isinstance(interruption, dict):
|
||||
heard = str(interruption.get("heard_text") or "").strip()
|
||||
current_text = str(message_text or heard).strip()
|
||||
if not last_interruptible and last_text:
|
||||
return InterruptionDecision(
|
||||
action="replay",
|
||||
text=current_text,
|
||||
replay_text=last_text,
|
||||
reason="non_interruptible_speech",
|
||||
is_interruptible=False,
|
||||
heard_text=heard,
|
||||
)
|
||||
if last_text:
|
||||
return InterruptionDecision(
|
||||
action="classify",
|
||||
text=current_text,
|
||||
replay_text=last_text,
|
||||
reason="interruptible_speech",
|
||||
is_interruptible=True,
|
||||
heard_text=heard,
|
||||
)
|
||||
return InterruptionDecision(
|
||||
action="process",
|
||||
text=current_text,
|
||||
reason="interruptible_speech_no_history",
|
||||
is_interruptible=True,
|
||||
heard_text=heard,
|
||||
)
|
||||
|
||||
return InterruptionDecision(action="process", text=message_text)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"InterruptionDecision",
|
||||
"classify_processing_interruption",
|
||||
"evaluate_interruption",
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
"""Correções determinísticas e conservadoras para transcrição de canal de voz."""
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Mapping
|
||||
|
||||
# Só falas inteiras entram nesta tabela. Nunca substitua tokens dentro de frases.
|
||||
DEFAULT_WHOLE_UTTERANCE_FIXES: dict[str, str] = {
|
||||
"fim": "Sim",
|
||||
"mim": "Sim",
|
||||
}
|
||||
|
||||
_TRAILING_PUNCT = re.compile(r"[.!?]+$")
|
||||
|
||||
|
||||
def fix_whole_utterance_transcription(
|
||||
text: str,
|
||||
*,
|
||||
fixes: Mapping[str, str] | None = None,
|
||||
) -> str:
|
||||
raw = str(text or "")
|
||||
stripped = raw.strip()
|
||||
if not stripped:
|
||||
return raw
|
||||
candidate = _TRAILING_PUNCT.sub("", stripped).strip().casefold()
|
||||
table = fixes or DEFAULT_WHOLE_UTTERANCE_FIXES
|
||||
replacement = table.get(candidate)
|
||||
return str(replacement) if replacement is not None else raw
|
||||
|
||||
|
||||
__all__ = ["DEFAULT_WHOLE_UTTERANCE_FIXES", "fix_whole_utterance_transcription"]
|
||||
@@ -0,0 +1,32 @@
|
||||
from .checkpoint_repository import (
|
||||
AutonomousCheckpointRepository,
|
||||
CheckpointIntegrityError,
|
||||
CheckpointIntegrityService,
|
||||
CheckpointRecoveryError,
|
||||
InMemoryCheckpointRepository,
|
||||
LangGraphCheckpointRepository,
|
||||
OracleCheckpointRepository,
|
||||
ResilientCheckpointRepository,
|
||||
RetryPolicy,
|
||||
SQLiteCheckpointRepository,
|
||||
create_checkpoint_repository,
|
||||
create_raw_checkpoint_repository,
|
||||
)
|
||||
from .langgraph_saver import RepositoryCheckpointSaver, create_langgraph_checkpointer
|
||||
|
||||
__all__ = [
|
||||
"AutonomousCheckpointRepository",
|
||||
"CheckpointIntegrityError",
|
||||
"CheckpointIntegrityService",
|
||||
"CheckpointRecoveryError",
|
||||
"InMemoryCheckpointRepository",
|
||||
"LangGraphCheckpointRepository",
|
||||
"OracleCheckpointRepository",
|
||||
"RepositoryCheckpointSaver",
|
||||
"ResilientCheckpointRepository",
|
||||
"RetryPolicy",
|
||||
"SQLiteCheckpointRepository",
|
||||
"create_checkpoint_repository",
|
||||
"create_langgraph_checkpointer",
|
||||
"create_raw_checkpoint_repository",
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,425 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import random
|
||||
import time
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Iterable
|
||||
|
||||
from agent_framework.persistence.sqlite_store import SQLiteStore
|
||||
|
||||
logger = logging.getLogger("agent_framework.checkpoints")
|
||||
|
||||
|
||||
def _utc_now() -> str:
|
||||
return datetime.now(timezone.utc).isoformat()
|
||||
|
||||
|
||||
def _json_dumps(value: Any) -> str:
|
||||
return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"), default=str)
|
||||
|
||||
|
||||
def _json_loads(value: str | bytes | None, default: Any):
|
||||
if value is None:
|
||||
return default
|
||||
if isinstance(value, bytes):
|
||||
value = value.decode("utf-8")
|
||||
try:
|
||||
return json.loads(value)
|
||||
except Exception:
|
||||
return default
|
||||
|
||||
|
||||
def _sha256(value: Any) -> str:
|
||||
return hashlib.sha256(_json_dumps(value).encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
class CheckpointIntegrityError(RuntimeError):
|
||||
"""Raised when a persisted checkpoint envelope fails checksum validation."""
|
||||
|
||||
|
||||
class CheckpointRecoveryError(RuntimeError):
|
||||
"""Raised when recovery cannot find a valid checkpoint."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RetryPolicy:
|
||||
max_attempts: int = 3
|
||||
base_delay_seconds: float = 0.05
|
||||
max_delay_seconds: float = 1.0
|
||||
jitter_seconds: float = 0.05
|
||||
|
||||
|
||||
class CheckpointIntegrityService:
|
||||
"""Creates and validates immutable checkpoint envelopes.
|
||||
|
||||
The repository stores an envelope instead of only the raw LangGraph payload:
|
||||
- schema_version: enables future migrations;
|
||||
- payload_hash: SHA-256 over the payload;
|
||||
- envelope_id: idempotency/correlation id;
|
||||
- compacted: marks synthetic compacted snapshots.
|
||||
"""
|
||||
|
||||
SCHEMA_VERSION = 1
|
||||
ENVELOPE_MARKER = "agent_framework_checkpoint_envelope"
|
||||
|
||||
def wrap(self, thread_id: str, checkpoint: dict[str, Any], *, compacted: bool = False) -> dict[str, Any]:
|
||||
payload = checkpoint or {}
|
||||
return {
|
||||
"_type": self.ENVELOPE_MARKER,
|
||||
"schema_version": self.SCHEMA_VERSION,
|
||||
"envelope_id": str(uuid.uuid4()),
|
||||
"thread_id": thread_id,
|
||||
"checkpoint_id": str(payload.get("checkpoint_id") or (payload.get("checkpoint") or {}).get("id") or uuid.uuid4()),
|
||||
"payload_hash": _sha256(payload),
|
||||
"payload": payload,
|
||||
"compacted": bool(compacted),
|
||||
"created_at": _utc_now(),
|
||||
}
|
||||
|
||||
def is_envelope(self, value: dict[str, Any] | None) -> bool:
|
||||
return isinstance(value, dict) and value.get("_type") == self.ENVELOPE_MARKER
|
||||
|
||||
def unwrap(self, value: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||
if value is None:
|
||||
return None
|
||||
if not self.is_envelope(value):
|
||||
# Backwards compatibility with old checkpoints from previous project versions.
|
||||
return value
|
||||
expected = value.get("payload_hash")
|
||||
payload = value.get("payload") or {}
|
||||
actual = _sha256(payload)
|
||||
if expected != actual:
|
||||
raise CheckpointIntegrityError(
|
||||
f"Checkpoint corrompido para thread_id={value.get('thread_id')}: hash esperado={expected}, hash atual={actual}"
|
||||
)
|
||||
if int(value.get("schema_version") or 0) > self.SCHEMA_VERSION:
|
||||
raise CheckpointIntegrityError(
|
||||
f"Checkpoint usa schema_version={value.get('schema_version')} maior que o suportado={self.SCHEMA_VERSION}"
|
||||
)
|
||||
return payload
|
||||
|
||||
|
||||
class LangGraphCheckpointRepository(ABC):
|
||||
@abstractmethod
|
||||
async def put(self, thread_id: str, checkpoint: dict[str, Any]) -> None: ...
|
||||
|
||||
@abstractmethod
|
||||
async def get_latest(self, thread_id: str) -> dict[str, Any] | None: ...
|
||||
|
||||
async def list_latest(self, thread_id: str, limit: int = 20) -> list[dict[str, Any]]:
|
||||
latest = await self.get_latest(thread_id)
|
||||
return [latest] if latest else []
|
||||
|
||||
async def compact(self, thread_id: str, keep_last: int = 20) -> int:
|
||||
return 0
|
||||
|
||||
@staticmethod
|
||||
def is_valid_checkpoint(checkpoint):
|
||||
if not isinstance(checkpoint, dict):
|
||||
return False
|
||||
if "v" in checkpoint:
|
||||
return True
|
||||
if (
|
||||
"checkpoint" in checkpoint
|
||||
and isinstance(checkpoint["checkpoint"], dict)
|
||||
and "v" in checkpoint["checkpoint"]
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
class InMemoryCheckpointRepository(LangGraphCheckpointRepository):
|
||||
def __init__(self):
|
||||
self._data: dict[str, list[dict[str, Any]]] = {}
|
||||
|
||||
async def put(self, thread_id: str, checkpoint: dict[str, Any]):
|
||||
self._data.setdefault(thread_id, []).append(checkpoint)
|
||||
|
||||
async def get_latest(self, thread_id: str):
|
||||
items = self._data.get(thread_id, [])
|
||||
return items[-1] if items else None
|
||||
|
||||
async def list_latest(self, thread_id: str, limit: int = 20) -> list[dict[str, Any]]:
|
||||
return list(reversed(self._data.get(thread_id, [])[-limit:]))
|
||||
|
||||
async def compact(self, thread_id: str, keep_last: int = 20) -> int:
|
||||
items = self._data.get(thread_id, [])
|
||||
if len(items) <= keep_last:
|
||||
return 0
|
||||
removed = len(items) - keep_last
|
||||
self._data[thread_id] = items[-keep_last:]
|
||||
return removed
|
||||
|
||||
|
||||
class SQLiteCheckpointRepository(LangGraphCheckpointRepository):
|
||||
def __init__(self, settings):
|
||||
self.store = SQLiteStore(settings.SQLITE_DB_PATH)
|
||||
|
||||
async def put(self, thread_id: str, checkpoint: dict[str, Any]):
|
||||
await asyncio.to_thread(self.store.put_checkpoint, thread_id, checkpoint)
|
||||
|
||||
async def get_latest(self, thread_id: str):
|
||||
return await asyncio.to_thread(self.store.get_latest_checkpoint, thread_id)
|
||||
|
||||
async def list_latest(self, thread_id: str, limit: int = 20) -> list[dict[str, Any]]:
|
||||
def _list():
|
||||
with self.store.connect() as con:
|
||||
rows = con.execute(
|
||||
"select checkpoint_json from workflow_checkpoints where thread_id=? order by id desc limit ?",
|
||||
(thread_id, int(limit)),
|
||||
).fetchall()
|
||||
return [_json_loads(r["checkpoint_json"], None) for r in rows if r]
|
||||
|
||||
return await asyncio.to_thread(_list)
|
||||
|
||||
async def compact(self, thread_id: str, keep_last: int = 20) -> int:
|
||||
def _compact():
|
||||
with self.store.connect() as con:
|
||||
rows = con.execute(
|
||||
"select id from workflow_checkpoints where thread_id=? order by id desc",
|
||||
(thread_id,),
|
||||
).fetchall()
|
||||
ids = [int(r["id"]) for r in rows]
|
||||
delete_ids = ids[int(keep_last):]
|
||||
if not delete_ids:
|
||||
return 0
|
||||
con.executemany("delete from workflow_checkpoints where id=?", [(i,) for i in delete_ids])
|
||||
return len(delete_ids)
|
||||
|
||||
return await asyncio.to_thread(_compact)
|
||||
|
||||
|
||||
class OracleCheckpointRepository(LangGraphCheckpointRepository):
|
||||
"""Checkpoint repository real para Oracle/Autonomous Database.
|
||||
|
||||
O OracleStore já cria as tabelas FIRST-compatible. A compactação é best-effort:
|
||||
remove checkpoints antigos quando o store expõe conexão e prefixo de tabelas.
|
||||
"""
|
||||
|
||||
def __init__(self, settings):
|
||||
from agent_framework.persistence.oracle_store import OracleStore
|
||||
|
||||
self.store = OracleStore(settings)
|
||||
|
||||
async def put(self, thread_id: str, checkpoint: dict[str, Any]):
|
||||
await self.store.put_checkpoint(thread_id, checkpoint)
|
||||
|
||||
async def get_latest(self, thread_id: str):
|
||||
return await self.store.get_latest_checkpoint(thread_id)
|
||||
|
||||
async def list_latest(self, thread_id: str, limit: int = 20) -> list[dict[str, Any]]:
|
||||
if not hasattr(self.store, "connect") or not hasattr(self.store, "t"):
|
||||
return await super().list_latest(thread_id, limit)
|
||||
|
||||
def _list():
|
||||
sql = f"""
|
||||
select CHECKPOINT_JSON
|
||||
from {self.store.t('WORKFLOW_CHECKPOINT')}
|
||||
where THREAD_ID = :thread_id
|
||||
order by ID desc
|
||||
fetch first :limit rows only
|
||||
"""
|
||||
with self.store.connect() as conn:
|
||||
rows = conn.cursor().execute(sql, dict(thread_id=thread_id, limit=int(limit))).fetchall()
|
||||
return [_json_loads(r[0], None) for r in rows if r]
|
||||
|
||||
return await asyncio.to_thread(_list)
|
||||
|
||||
async def compact(self, thread_id: str, keep_last: int = 20) -> int:
|
||||
if not hasattr(self.store, "connect") or not hasattr(self.store, "t"):
|
||||
return 0
|
||||
|
||||
def _compact():
|
||||
table = self.store.t("WORKFLOW_CHECKPOINT")
|
||||
sql_count = f"select count(*) from {table} where THREAD_ID = :thread_id"
|
||||
sql_delete = f"""
|
||||
delete from {table}
|
||||
where THREAD_ID = :thread_id
|
||||
and ID not in (
|
||||
select ID from {table}
|
||||
where THREAD_ID = :thread_id
|
||||
order by ID desc
|
||||
fetch first :keep_last rows only
|
||||
)
|
||||
"""
|
||||
with self.store.connect() as conn:
|
||||
cur = conn.cursor()
|
||||
before = int(cur.execute(sql_count, dict(thread_id=thread_id)).fetchone()[0])
|
||||
cur.execute(sql_delete, dict(thread_id=thread_id, keep_last=int(keep_last)))
|
||||
after = int(cur.execute(sql_count, dict(thread_id=thread_id)).fetchone()[0])
|
||||
return max(0, before - after)
|
||||
|
||||
return await asyncio.to_thread(_compact)
|
||||
|
||||
|
||||
AutonomousCheckpointRepository = OracleCheckpointRepository
|
||||
|
||||
|
||||
class ResilientCheckpointRepository(LangGraphCheckpointRepository):
|
||||
"""Adds integrity, retry, compaction and recovery to any repository.
|
||||
|
||||
This wrapper is intentionally repository-neutral. It can protect memory,
|
||||
SQLite and Oracle repositories without changing LangGraph code.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
inner: LangGraphCheckpointRepository,
|
||||
*,
|
||||
integrity: CheckpointIntegrityService | None = None,
|
||||
retry_policy: RetryPolicy | None = None,
|
||||
enable_integrity: bool = True,
|
||||
enable_compaction: bool = True,
|
||||
compact_every: int = 50,
|
||||
keep_last: int = 20,
|
||||
recovery_scan_limit: int = 25,
|
||||
):
|
||||
self.inner = inner
|
||||
self.integrity = integrity or CheckpointIntegrityService()
|
||||
self.retry_policy = retry_policy or RetryPolicy()
|
||||
self.enable_integrity = enable_integrity
|
||||
self.enable_compaction = enable_compaction
|
||||
self.compact_every = max(1, int(compact_every))
|
||||
self.keep_last = max(1, int(keep_last))
|
||||
self.recovery_scan_limit = max(1, int(recovery_scan_limit))
|
||||
self._put_count_by_thread: dict[str, int] = {}
|
||||
|
||||
async def _with_retry(self, operation_name: str, coro_factory):
|
||||
last_exc: Exception | None = None
|
||||
for attempt in range(1, self.retry_policy.max_attempts + 1):
|
||||
try:
|
||||
return await coro_factory()
|
||||
except Exception as exc: # noqa: BLE001 - repository failures vary by backend
|
||||
last_exc = exc
|
||||
if attempt >= self.retry_policy.max_attempts:
|
||||
break
|
||||
delay = min(
|
||||
self.retry_policy.max_delay_seconds,
|
||||
self.retry_policy.base_delay_seconds * (2 ** (attempt - 1)),
|
||||
) + random.uniform(0, self.retry_policy.jitter_seconds)
|
||||
logger.warning("checkpoint.%s.retry attempt=%s delay=%.3fs error=%s", operation_name, attempt, delay, exc)
|
||||
await asyncio.sleep(delay)
|
||||
raise last_exc # type: ignore[misc]
|
||||
|
||||
async def put(self, thread_id: str, checkpoint: dict[str, Any]) -> None:
|
||||
payload = self.integrity.wrap(thread_id, checkpoint) if self.enable_integrity else checkpoint
|
||||
await self._with_retry("put", lambda: self.inner.put(thread_id, payload))
|
||||
self._put_count_by_thread[thread_id] = self._put_count_by_thread.get(thread_id, 0) + 1
|
||||
if self.enable_compaction and self._put_count_by_thread[thread_id] % self.compact_every == 0:
|
||||
try:
|
||||
removed = await self.inner.compact(thread_id, keep_last=self.keep_last)
|
||||
if removed:
|
||||
logger.info("checkpoint.compaction thread_id=%s removed=%s keep_last=%s", thread_id, removed, self.keep_last)
|
||||
except Exception as exc: # compaction must never break the user flow
|
||||
logger.warning("checkpoint.compaction.failed thread_id=%s error=%s", thread_id, exc)
|
||||
|
||||
async def get_latest(self, thread_id: str) -> dict[str, Any] | None:
|
||||
return await self.recover_latest(thread_id)
|
||||
|
||||
async def list_latest(self, thread_id: str, limit: int = 20) -> list[dict[str, Any]]:
|
||||
raw_items = await self.inner.list_latest(thread_id, limit)
|
||||
out: list[dict[str, Any]] = []
|
||||
for item in raw_items:
|
||||
try:
|
||||
payload = self.integrity.unwrap(item) if self.enable_integrity else item
|
||||
if payload is not None:
|
||||
out.append(payload)
|
||||
except CheckpointIntegrityError:
|
||||
continue
|
||||
return out
|
||||
|
||||
async def compact(self, thread_id: str, keep_last: int = 20) -> int:
|
||||
return await self.inner.compact(thread_id, keep_last=keep_last)
|
||||
|
||||
async def recover_latest(self, thread_id: str) -> dict[str, Any] | None:
|
||||
"""Return the newest valid LangGraph checkpoint, skipping corrupt or legacy records."""
|
||||
raw_items = await self._with_retry(
|
||||
"list_latest",
|
||||
lambda: self.inner.list_latest(thread_id, self.recovery_scan_limit),
|
||||
)
|
||||
|
||||
first_integrity_error: Exception | None = None
|
||||
invalid_count = 0
|
||||
|
||||
for raw in raw_items:
|
||||
try:
|
||||
payload = self.integrity.unwrap(raw)
|
||||
|
||||
candidate = payload
|
||||
|
||||
if (
|
||||
isinstance(payload, dict)
|
||||
and "checkpoint" in payload
|
||||
):
|
||||
candidate = payload["checkpoint"]
|
||||
|
||||
if not self.is_valid_checkpoint(candidate):
|
||||
continue
|
||||
|
||||
return payload
|
||||
|
||||
except CheckpointIntegrityError as exc:
|
||||
first_integrity_error = first_integrity_error or exc
|
||||
logger.error(
|
||||
"checkpoint.recovery.skip_corrupt thread_id=%s error=%s",
|
||||
thread_id,
|
||||
exc,
|
||||
)
|
||||
continue
|
||||
|
||||
if first_integrity_error:
|
||||
# No valid checkpoint: return None so the run starts clean instead of crashing ainvoke.
|
||||
logger.error(
|
||||
"checkpoint.recovery.no_valid_checkpoint thread_id=%s starting_fresh error=%s",
|
||||
thread_id,
|
||||
first_integrity_error,
|
||||
)
|
||||
return None
|
||||
|
||||
if invalid_count:
|
||||
logger.warning(
|
||||
"checkpoint.recovery.no_valid_langgraph_checkpoint "
|
||||
"thread_id=%s invalid_count=%s",
|
||||
thread_id,
|
||||
invalid_count,
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def _retry_policy_from_settings(settings) -> RetryPolicy:
|
||||
return RetryPolicy(
|
||||
max_attempts=int(getattr(settings, "CHECKPOINT_RETRY_MAX_ATTEMPTS", 3) or 3),
|
||||
base_delay_seconds=float(getattr(settings, "CHECKPOINT_RETRY_BASE_DELAY_SECONDS", 0.05) or 0.05),
|
||||
max_delay_seconds=float(getattr(settings, "CHECKPOINT_RETRY_MAX_DELAY_SECONDS", 1.0) or 1.0),
|
||||
jitter_seconds=float(getattr(settings, "CHECKPOINT_RETRY_JITTER_SECONDS", 0.05) or 0.05),
|
||||
)
|
||||
|
||||
|
||||
def create_raw_checkpoint_repository(settings):
|
||||
provider = getattr(settings, "CHECKPOINT_REPOSITORY_PROVIDER", "memory")
|
||||
if provider == "sqlite":
|
||||
return SQLiteCheckpointRepository(settings)
|
||||
if provider in {"autonomous", "oracle"}:
|
||||
return OracleCheckpointRepository(settings)
|
||||
return InMemoryCheckpointRepository()
|
||||
|
||||
|
||||
def create_checkpoint_repository(settings):
|
||||
raw = create_raw_checkpoint_repository(settings)
|
||||
if not bool(getattr(settings, "ENABLE_RESILIENT_CHECKPOINTER", True)):
|
||||
return raw
|
||||
return ResilientCheckpointRepository(
|
||||
raw,
|
||||
retry_policy=_retry_policy_from_settings(settings),
|
||||
enable_integrity=bool(getattr(settings, "ENABLE_CHECKPOINT_INTEGRITY", True)),
|
||||
enable_compaction=bool(getattr(settings, "ENABLE_CHECKPOINT_COMPACTION", True)),
|
||||
compact_every=int(getattr(settings, "CHECKPOINT_COMPACT_EVERY", 50) or 50),
|
||||
keep_last=int(getattr(settings, "CHECKPOINT_KEEP_LAST", 20) or 20),
|
||||
recovery_scan_limit=int(getattr(settings, "CHECKPOINT_RECOVERY_SCAN_LIMIT", 25) or 25),
|
||||
)
|
||||
@@ -0,0 +1,454 @@
|
||||
from __future__ import annotations
|
||||
try:
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
except Exception: # pragma: no cover - fallback for lightweight unit tests without langgraph installed
|
||||
class BaseCheckpointSaver: # type: ignore[no-redef]
|
||||
pass
|
||||
|
||||
"""LangGraph checkpoint saver backed by the framework checkpoint repository.
|
||||
|
||||
This module intentionally keeps a small adapter surface so the framework can run
|
||||
with multiple LangGraph versions. It implements the common synchronous and
|
||||
asynchronous methods used by BaseCheckpointSaver/MemorySaver: get_tuple,
|
||||
aget_tuple, put, aput, put_writes, aput_writes, list and alist.
|
||||
|
||||
The persisted payload stores LangGraph's raw checkpoint/config/metadata values in
|
||||
repository-neutral JSON. When LangGraph is installed, checkpoint tuples are
|
||||
returned using CheckpointTuple; otherwise a simple dict is returned for tests.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import uuid
|
||||
from typing import Any, AsyncIterator, Iterator
|
||||
|
||||
from .checkpoint_repository import create_checkpoint_repository
|
||||
|
||||
|
||||
def _parse_legacy_json_container(value: Any, expected: type) -> Any:
|
||||
"""Recover containers that older JSON backends persisted as JSON strings.
|
||||
|
||||
This is intentionally field-scoped: ordinary business strings must stay
|
||||
strings, even if their text happens to look like JSON.
|
||||
"""
|
||||
current = value
|
||||
for _ in range(3):
|
||||
if isinstance(current, expected):
|
||||
return current
|
||||
if not isinstance(current, str):
|
||||
break
|
||||
text = current.strip()
|
||||
if not text:
|
||||
break
|
||||
if expected is dict and not text.startswith("{"):
|
||||
break
|
||||
if expected is list and not text.startswith("["):
|
||||
break
|
||||
try:
|
||||
current = json.loads(text)
|
||||
except Exception:
|
||||
break
|
||||
return current if isinstance(current, expected) else expected()
|
||||
|
||||
|
||||
def _strict_json_value(value: Any, *, path: str = "$") -> Any:
|
||||
"""Convert to repository-safe JSON without ever falling back to ``str``.
|
||||
|
||||
``default=str`` is unsafe for LangGraph checkpoints: runtime/task objects can
|
||||
become ordinary strings and later be consumed as typed values by Pregel.
|
||||
Keep native JSON containers recursively and fail loudly for an unsupported
|
||||
object instead of corrupting it silently.
|
||||
"""
|
||||
if value is None or isinstance(value, (str, int, float, bool)):
|
||||
return value
|
||||
if isinstance(value, dict):
|
||||
return {
|
||||
str(key): _strict_json_value(item, path=f"{path}.{key}")
|
||||
for key, item in value.items()
|
||||
}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [
|
||||
_strict_json_value(item, path=f"{path}[{idx}]")
|
||||
for idx, item in enumerate(value)
|
||||
]
|
||||
# Common durable scalar types that JSON does not know natively.
|
||||
if isinstance(value, uuid.UUID):
|
||||
return str(value)
|
||||
try:
|
||||
from datetime import date, datetime
|
||||
if isinstance(value, (date, datetime)):
|
||||
return value.isoformat()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
from enum import Enum
|
||||
if isinstance(value, Enum):
|
||||
return _strict_json_value(value.value, path=path)
|
||||
except Exception:
|
||||
pass
|
||||
if hasattr(value, "model_dump") and callable(value.model_dump):
|
||||
return _strict_json_value(value.model_dump(), path=path)
|
||||
raise TypeError(
|
||||
f"Checkpoint contém valor não serializável em {path}: "
|
||||
f"{type(value).__module__}.{type(value).__qualname__}"
|
||||
)
|
||||
|
||||
|
||||
def _normalize_checkpoint(checkpoint: Any) -> dict[str, Any]:
|
||||
checkpoint = _parse_legacy_json_container(checkpoint, dict)
|
||||
if not isinstance(checkpoint, dict):
|
||||
return {}
|
||||
out = dict(checkpoint)
|
||||
out["channel_values"] = _parse_legacy_json_container(out.get("channel_values"), dict)
|
||||
out["channel_versions"] = _parse_legacy_json_container(out.get("channel_versions"), dict)
|
||||
raw_seen = _parse_legacy_json_container(out.get("versions_seen"), dict)
|
||||
out["versions_seen"] = {
|
||||
str(node): _parse_legacy_json_container(versions, dict)
|
||||
for node, versions in raw_seen.items()
|
||||
}
|
||||
if "pending_sends" in out:
|
||||
out["pending_sends"] = _parse_legacy_json_container(out.get("pending_sends"), list)
|
||||
if "updated_channels" in out and isinstance(out.get("updated_channels"), str):
|
||||
out["updated_channels"] = _parse_legacy_json_container(out.get("updated_channels"), list)
|
||||
return out
|
||||
|
||||
|
||||
def _normalize_metadata(metadata: Any) -> dict[str, Any]:
|
||||
value = _parse_legacy_json_container(metadata, dict)
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
||||
|
||||
def _normalize_config(config: Any) -> dict[str, Any]:
|
||||
value = _parse_legacy_json_container(config, dict)
|
||||
if not isinstance(value, dict):
|
||||
return {}
|
||||
out = dict(value)
|
||||
out["configurable"] = _parse_legacy_json_container(out.get("configurable"), dict)
|
||||
return out
|
||||
|
||||
|
||||
_EPHEMERAL_RUNTIME_KEYS = {"__pregel_runtime", "__pregel_store"}
|
||||
|
||||
|
||||
def _strip_runtime_refs(value: Any) -> Any:
|
||||
"""Recursively remove process-local runtime/store references only.
|
||||
|
||||
Checkpoints may legitimately contain LangGraph internal channels whose names
|
||||
also start with ``__pregel_`` (for example task channels). Those are durable
|
||||
graph state and must be preserved. The corruption that triggers
|
||||
``str.override`` is specifically a runtime/store object captured inside a
|
||||
nested RunnableConfig and later stringified by the JSON repository.
|
||||
"""
|
||||
if isinstance(value, dict):
|
||||
return {
|
||||
key: _strip_runtime_refs(item)
|
||||
for key, item in value.items()
|
||||
if str(key) not in _EPHEMERAL_RUNTIME_KEYS
|
||||
}
|
||||
if isinstance(value, list):
|
||||
return [_strip_runtime_refs(item) for item in value]
|
||||
if isinstance(value, tuple):
|
||||
return tuple(_strip_runtime_refs(item) for item in value)
|
||||
return value
|
||||
|
||||
|
||||
def _durable_config(config: dict[str, Any] | None) -> dict[str, Any]:
|
||||
"""Return a checkpoint-safe copy of a LangGraph RunnableConfig.
|
||||
|
||||
LangGraph injects ephemeral private values such as ``__pregel_runtime`` and
|
||||
``__pregel_store`` under ``configurable`` while a graph is running. They are
|
||||
process-local and must never cross the durable checkpoint boundary.
|
||||
|
||||
The scrub is recursive because task/pending-write config fragments may be
|
||||
nested below regular config fields in newer LangGraph versions.
|
||||
"""
|
||||
if not isinstance(config, dict):
|
||||
return {}
|
||||
cleaned = _strip_runtime_refs(config)
|
||||
if not isinstance(cleaned, dict):
|
||||
return {}
|
||||
configurable = cleaned.get("configurable")
|
||||
if isinstance(configurable, dict):
|
||||
cleaned = dict(cleaned)
|
||||
cleaned["configurable"] = {
|
||||
key: value
|
||||
for key, value in configurable.items()
|
||||
if not str(key).startswith("__pregel_")
|
||||
}
|
||||
return cleaned
|
||||
|
||||
|
||||
def _canonical_checkpoint_config(
|
||||
payload: dict[str, Any],
|
||||
request_config: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Rebuild the RunnableConfig returned to LangGraph from durable IDs only.
|
||||
|
||||
Official LangGraph savers do not re-bind the full config that happened to be
|
||||
present when a checkpoint was written. They reconstruct a fresh config from
|
||||
``thread_id``, ``checkpoint_ns`` and ``checkpoint_id``. Doing the same here
|
||||
prevents a historical/factory-time runtime value from being rebound into a
|
||||
new execution while remaining backward compatible with existing rows.
|
||||
"""
|
||||
requested = _durable_config(request_config)
|
||||
stored = _durable_config(_normalize_config(payload.get("config")) if isinstance(payload, dict) else None)
|
||||
req_cfg = requested.get("configurable") if isinstance(requested.get("configurable"), dict) else {}
|
||||
stored_cfg = stored.get("configurable") if isinstance(stored.get("configurable"), dict) else {}
|
||||
checkpoint = payload.get("checkpoint") if isinstance(payload, dict) else {}
|
||||
checkpoint = checkpoint if isinstance(checkpoint, dict) else {}
|
||||
|
||||
thread_id = (
|
||||
req_cfg.get("thread_id")
|
||||
or stored_cfg.get("thread_id")
|
||||
or payload.get("thread_id")
|
||||
or "default"
|
||||
)
|
||||
checkpoint_ns = req_cfg.get("checkpoint_ns")
|
||||
if checkpoint_ns is None:
|
||||
checkpoint_ns = stored_cfg.get("checkpoint_ns", "")
|
||||
|
||||
requested_checkpoint_id = req_cfg.get("checkpoint_id")
|
||||
checkpoint_id = (
|
||||
requested_checkpoint_id
|
||||
or payload.get("checkpoint_id")
|
||||
or checkpoint.get("id")
|
||||
or stored_cfg.get("checkpoint_id")
|
||||
)
|
||||
|
||||
configurable: dict[str, Any] = {
|
||||
"thread_id": str(thread_id),
|
||||
"checkpoint_ns": str(checkpoint_ns or ""),
|
||||
}
|
||||
if checkpoint_id not in (None, ""):
|
||||
configurable["checkpoint_id"] = str(checkpoint_id)
|
||||
return {"configurable": configurable}
|
||||
|
||||
|
||||
def _thread_id(config: dict[str, Any] | None) -> str:
|
||||
configurable = (config or {}).get("configurable") or {}
|
||||
return str(configurable.get("thread_id") or configurable.get("checkpoint_ns") or "default")
|
||||
|
||||
|
||||
def _checkpoint_id(checkpoint: dict[str, Any] | None) -> str:
|
||||
if isinstance(checkpoint, dict):
|
||||
return str(checkpoint.get("id") or checkpoint.get("checkpoint_id") or uuid.uuid4())
|
||||
return str(uuid.uuid4())
|
||||
|
||||
|
||||
def _normalize_pending_writes(pending_writes: Any) -> list[tuple[Any, Any, Any]]:
|
||||
"""Normalize persisted pending_writes to LangGraph's expected runtime format.
|
||||
|
||||
LangGraph 1.1.x expects CheckpointTuple.pending_writes to be an iterable of
|
||||
3-item tuples: (task_id, channel, value).
|
||||
|
||||
Older framework versions persisted writes as dictionaries containing
|
||||
task_id, task_path, channel and value. Some stores/tests may also contain
|
||||
4-item tuples: (task_id, task_path, channel, value). This adapter accepts
|
||||
those legacy forms while preserving already-correct 3-item tuples.
|
||||
"""
|
||||
normalized: list[tuple[Any, Any, Any]] = []
|
||||
for item in pending_writes or []:
|
||||
if isinstance(item, dict):
|
||||
normalized.append((
|
||||
item.get("task_id"),
|
||||
item.get("channel"),
|
||||
item.get("value"),
|
||||
))
|
||||
continue
|
||||
|
||||
if isinstance(item, (list, tuple)):
|
||||
if len(item) == 3:
|
||||
task_id, channel, value = item
|
||||
normalized.append((task_id, channel, value))
|
||||
continue
|
||||
if len(item) == 4:
|
||||
task_id, _task_path, channel, value = item
|
||||
normalized.append((task_id, channel, value))
|
||||
continue
|
||||
|
||||
# Defensive fallback: keep malformed legacy entries from crashing resume.
|
||||
# Use a synthetic channel so the data remains inspectable in telemetry/logs.
|
||||
normalized.append((None, "__malformed_pending_write__", item))
|
||||
return normalized
|
||||
|
||||
|
||||
class RepositoryCheckpointSaver(BaseCheckpointSaver):
|
||||
"""Checkpoint saver nativo para LangGraph usando os repositories do framework."""
|
||||
|
||||
def __init__(self, settings, repository=None):
|
||||
super().__init__()
|
||||
self.settings = settings
|
||||
self.repository = repository or create_checkpoint_repository(settings)
|
||||
self._loop: asyncio.AbstractEventLoop | None = None
|
||||
|
||||
def _run(self, coro):
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return asyncio.run(coro)
|
||||
# LangGraph may call sync methods from a worker thread; when already in
|
||||
# an event loop prefer a short-lived thread to avoid nested-loop errors.
|
||||
import concurrent.futures
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as ex:
|
||||
return ex.submit(lambda: asyncio.run(coro)).result()
|
||||
|
||||
def _make_tuple(
|
||||
self,
|
||||
payload: dict[str, Any] | None,
|
||||
request_config: dict[str, Any] | None = None,
|
||||
):
|
||||
if not payload:
|
||||
return None
|
||||
# Second-stage protection: never re-bind the full persisted RunnableConfig.
|
||||
# Rebuild only the durable identifiers, as official LangGraph savers do.
|
||||
config = _canonical_checkpoint_config(payload, request_config)
|
||||
checkpoint = _strip_runtime_refs(_normalize_checkpoint(payload.get("checkpoint") or {}))
|
||||
metadata = _strip_runtime_refs(_normalize_metadata(payload.get("metadata") or {}))
|
||||
raw_parent_config = payload.get("parent_config")
|
||||
if isinstance(raw_parent_config, dict):
|
||||
parent_payload = {
|
||||
"thread_id": payload.get("thread_id"),
|
||||
"config": raw_parent_config,
|
||||
"checkpoint_id": (raw_parent_config.get("configurable") or {}).get("checkpoint_id")
|
||||
if isinstance(raw_parent_config.get("configurable"), dict)
|
||||
else None,
|
||||
"checkpoint": {},
|
||||
}
|
||||
parent_config = _canonical_checkpoint_config(parent_payload)
|
||||
else:
|
||||
parent_config = None
|
||||
pending_writes = _normalize_pending_writes(
|
||||
_strip_runtime_refs(payload.get("pending_writes") or [])
|
||||
)
|
||||
try:
|
||||
from langgraph.checkpoint.base import CheckpointTuple
|
||||
return CheckpointTuple(config=config, checkpoint=checkpoint, metadata=metadata, parent_config=parent_config, pending_writes=pending_writes)
|
||||
except Exception:
|
||||
return {
|
||||
"config": _durable_config(config),
|
||||
"checkpoint": checkpoint,
|
||||
"metadata": metadata,
|
||||
"parent_config": parent_config,
|
||||
"pending_writes": pending_writes,
|
||||
}
|
||||
|
||||
async def aget_tuple(self, config: dict[str, Any]):
|
||||
return self._make_tuple(
|
||||
await self.repository.get_latest(_thread_id(config)),
|
||||
request_config=config,
|
||||
)
|
||||
|
||||
def get_tuple(self, config: dict[str, Any]):
|
||||
return self._run(self.aget_tuple(config))
|
||||
|
||||
async def aput(self, config: dict[str, Any], checkpoint: dict[str, Any], metadata: dict[str, Any] | None = None, new_versions: dict[str, Any] | None = None):
|
||||
thread_id = _thread_id(config)
|
||||
checkpoint_id = _checkpoint_id(checkpoint)
|
||||
clean_config = _durable_config(config)
|
||||
clean_cfg = clean_config.get("configurable") if isinstance(clean_config.get("configurable"), dict) else {}
|
||||
checkpoint_ns = str(clean_cfg.get("checkpoint_ns") or "")
|
||||
# Return a fresh canonical config. Never feed process-local/factory-time
|
||||
# configurable values back into the next LangGraph super-step.
|
||||
next_config = {
|
||||
"configurable": {
|
||||
"thread_id": thread_id,
|
||||
"checkpoint_ns": checkpoint_ns,
|
||||
"checkpoint_id": checkpoint_id,
|
||||
}
|
||||
}
|
||||
await self.repository.put(thread_id, {
|
||||
"thread_id": thread_id,
|
||||
"config": _strict_json_value(next_config, path="$.config"),
|
||||
"checkpoint": _strict_json_value(_strip_runtime_refs(_normalize_checkpoint(checkpoint)), path="$.checkpoint"),
|
||||
"metadata": _strict_json_value(_strip_runtime_refs(_normalize_metadata(metadata or {})), path="$.metadata"),
|
||||
"new_versions": _strict_json_value(_strip_runtime_refs(new_versions or {}), path="$.new_versions"),
|
||||
"checkpoint_id": checkpoint_id,
|
||||
})
|
||||
return next_config
|
||||
|
||||
def put(self, config: dict[str, Any], checkpoint: dict[str, Any], metadata: dict[str, Any] | None = None, new_versions: dict[str, Any] | None = None):
|
||||
return self._run(self.aput(config, checkpoint, metadata, new_versions))
|
||||
|
||||
async def aput_writes(self, config: dict[str, Any], writes: list[tuple[str, Any]], task_id: str, task_path: str = ""):
|
||||
thread_id = _thread_id(config)
|
||||
try:
|
||||
latest = await self.repository.get_latest(thread_id) or {"thread_id": thread_id, "config": _durable_config(config), "checkpoint": {}, "metadata": {}}
|
||||
except:
|
||||
latest = {
|
||||
"thread_id": thread_id,
|
||||
"config": _durable_config(config),
|
||||
"checkpoint": {},
|
||||
"metadata": {},
|
||||
"pending_writes": [],
|
||||
}
|
||||
|
||||
if isinstance(latest, dict):
|
||||
# Do not keep extending a persisted RunnableConfig across super-steps.
|
||||
# Rebuild the same canonical config that aget_tuple() will expose.
|
||||
latest["config"] = _canonical_checkpoint_config(latest, config)
|
||||
if isinstance(latest.get("checkpoint"), dict):
|
||||
latest["checkpoint"] = _strip_runtime_refs(latest.get("checkpoint"))
|
||||
if isinstance(latest.get("metadata"), dict):
|
||||
latest["metadata"] = _strip_runtime_refs(latest.get("metadata"))
|
||||
if isinstance(latest.get("parent_config"), dict):
|
||||
parent_payload = {
|
||||
"thread_id": latest.get("thread_id") or thread_id,
|
||||
"config": latest.get("parent_config"),
|
||||
"checkpoint_id": (latest.get("parent_config", {}).get("configurable") or {}).get("checkpoint_id")
|
||||
if isinstance(latest.get("parent_config", {}).get("configurable"), dict)
|
||||
else None,
|
||||
"checkpoint": {},
|
||||
}
|
||||
latest["parent_config"] = _canonical_checkpoint_config(parent_payload)
|
||||
|
||||
pending = list(latest.get("pending_writes") or [])
|
||||
for channel, value in writes or []:
|
||||
# Writes may contain nested task/RunnableConfig fragments. Scrub the
|
||||
# private runtime before the repository's JSON ``default=str`` layer.
|
||||
durable_value = _strip_runtime_refs(value)
|
||||
pending.append({
|
||||
"task_id": task_id,
|
||||
"task_path": task_path,
|
||||
"channel": channel,
|
||||
"value": _strict_json_value(durable_value, path=f"$.pending_writes[{task_id}].{channel}"),
|
||||
})
|
||||
latest["pending_writes"] = pending
|
||||
await self.repository.put(thread_id, latest)
|
||||
|
||||
def put_writes(self, config: dict[str, Any], writes: list[tuple[str, Any]], task_id: str, task_path: str = ""):
|
||||
return self._run(self.aput_writes(config, writes, task_id, task_path))
|
||||
|
||||
async def alist(self, config: dict[str, Any] | None = None, *, filter: dict[str, Any] | None = None, before: dict[str, Any] | None = None, limit: int | None = None) -> AsyncIterator[Any]:
|
||||
# Repository interface currently exposes only latest; this is enough for
|
||||
# resume/recovery. Oracle/SQLite repositories can later implement full list.
|
||||
if config is None:
|
||||
return
|
||||
item = await self.aget_tuple(config)
|
||||
if item:
|
||||
yield item
|
||||
|
||||
def list(self, config: dict[str, Any] | None = None, *, filter: dict[str, Any] | None = None, before: dict[str, Any] | None = None, limit: int | None = None) -> Iterator[Any]:
|
||||
item = self.get_tuple(config or {}) if config else None
|
||||
if item:
|
||||
yield item
|
||||
|
||||
|
||||
def create_langgraph_checkpointer(settings):
|
||||
"""Factory used by applications when compiling LangGraph.
|
||||
|
||||
By default the framework now returns RepositoryCheckpointSaver even for
|
||||
CHECKPOINT_REPOSITORY_PROVIDER=memory, because the repository wrapper adds
|
||||
integrity checks, retry, recovery and compaction.
|
||||
|
||||
Set ENABLE_RESILIENT_CHECKPOINTER=false to fall back to LangGraph MemorySaver
|
||||
for very small local experiments.
|
||||
"""
|
||||
provider = getattr(settings, "CHECKPOINT_REPOSITORY_PROVIDER", "memory")
|
||||
resilient = bool(getattr(settings, "ENABLE_RESILIENT_CHECKPOINTER", True))
|
||||
if provider == "memory" and not resilient:
|
||||
try:
|
||||
from langgraph.checkpoint.memory import MemorySaver
|
||||
return MemorySaver()
|
||||
except Exception:
|
||||
return RepositoryCheckpointSaver(settings)
|
||||
return RepositoryCheckpointSaver(settings)
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,90 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
try:
|
||||
import yaml
|
||||
except Exception: # pragma: no cover
|
||||
yaml = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class AgentProfile:
|
||||
agent_id: str
|
||||
name: str = ""
|
||||
description: str = ""
|
||||
prompt_policy_path: str | None = None
|
||||
routing_config_path: str | None = None
|
||||
guardrails_config_path: str | None = None
|
||||
judges_config_path: str | None = None
|
||||
mcp_servers_config_path: str | None = None
|
||||
tools_config_path: str | None = None
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
class AgentProfileRegistry:
|
||||
"""Carrega perfis de agentes/templates a partir de YAML.
|
||||
|
||||
O objetivo é permitir múltiplos agent_template no mesmo backend sem misturar
|
||||
memória, checkpoints, prompts, guardrails ou judges.
|
||||
"""
|
||||
|
||||
def __init__(self, settings):
|
||||
self.settings = settings
|
||||
self.base_dir = Path.cwd()
|
||||
self.profiles: dict[str, AgentProfile] = {}
|
||||
self.default_agent_id = "default_agent"
|
||||
self._load()
|
||||
|
||||
def _resolve(self, value: str | None) -> str | None:
|
||||
if not value:
|
||||
return None
|
||||
path = Path(value)
|
||||
return str(path if path.is_absolute() else (self.base_dir / path).resolve())
|
||||
|
||||
def _load(self) -> None:
|
||||
config_path = Path(getattr(self.settings, "AGENTS_CONFIG_PATH", "./config/agents.yaml"))
|
||||
if not config_path.is_absolute():
|
||||
config_path = self.base_dir / config_path
|
||||
if not config_path.exists() or yaml is None:
|
||||
self.profiles[self.default_agent_id] = AgentProfile(
|
||||
agent_id=self.default_agent_id,
|
||||
name="Default Agent",
|
||||
prompt_policy_path=self._resolve(getattr(self.settings, "PROMPT_POLICY_PATH", None)),
|
||||
routing_config_path=self._resolve(getattr(self.settings, "ROUTING_CONFIG_PATH", None)),
|
||||
guardrails_config_path=self._resolve(getattr(self.settings, "GUARDRAILS_CONFIG_PATH", None)),
|
||||
judges_config_path=self._resolve(getattr(self.settings, "JUDGES_CONFIG_PATH", None)),
|
||||
mcp_servers_config_path=self._resolve(getattr(self.settings, "MCP_SERVERS_CONFIG_PATH", None)),
|
||||
tools_config_path=self._resolve(getattr(self.settings, "TOOLS_CONFIG_PATH", None)),
|
||||
)
|
||||
return
|
||||
|
||||
raw = yaml.safe_load(config_path.read_text(encoding="utf-8")) or {}
|
||||
self.default_agent_id = raw.get("default_agent_id") or self.default_agent_id
|
||||
for item in raw.get("agents", []):
|
||||
agent_id = str(item.get("agent_id") or item.get("id") or "").strip()
|
||||
if not agent_id:
|
||||
continue
|
||||
self.profiles[agent_id] = AgentProfile(
|
||||
agent_id=agent_id,
|
||||
name=item.get("name", agent_id),
|
||||
description=item.get("description", ""),
|
||||
prompt_policy_path=self._resolve(item.get("prompt_policy_path") or getattr(self.settings, "PROMPT_POLICY_PATH", None)),
|
||||
routing_config_path=self._resolve(item.get("routing_config_path") or getattr(self.settings, "ROUTING_CONFIG_PATH", None)),
|
||||
guardrails_config_path=self._resolve(item.get("guardrails_config_path") or getattr(self.settings, "GUARDRAILS_CONFIG_PATH", None)),
|
||||
judges_config_path=self._resolve(item.get("judges_config_path") or getattr(self.settings, "JUDGES_CONFIG_PATH", None)),
|
||||
mcp_servers_config_path=self._resolve(item.get("mcp_servers_config_path") or getattr(self.settings, "MCP_SERVERS_CONFIG_PATH", None)),
|
||||
tools_config_path=self._resolve(item.get("tools_config_path") or getattr(self.settings, "TOOLS_CONFIG_PATH", None)),
|
||||
metadata=item.get("metadata") or {},
|
||||
)
|
||||
if self.default_agent_id not in self.profiles and self.profiles:
|
||||
self.default_agent_id = next(iter(self.profiles))
|
||||
|
||||
def get(self, agent_id: str | None = None) -> AgentProfile:
|
||||
key = agent_id or self.default_agent_id
|
||||
return self.profiles.get(key) or self.profiles[self.default_agent_id]
|
||||
|
||||
def list_profiles(self) -> list[AgentProfile]:
|
||||
return list(self.profiles.values())
|
||||
@@ -0,0 +1,82 @@
|
||||
version: "2"
|
||||
|
||||
# Default compatibility registry shipped with agent_framework_oci.
|
||||
#
|
||||
# This file reproduces the historical behavior that used to be hardcoded in
|
||||
# OutputSupervisor / ParallelRailExecutor. It is ALWAYS loaded by the framework.
|
||||
# An agent/deployment observability_mapping.yaml is then applied as an overlay.
|
||||
#
|
||||
# Therefore an older agent can replace only the framework and keep the same
|
||||
# GRL contract and legacy guardrail actions without adding new configuration.
|
||||
mappings:
|
||||
# Historical OutputSupervisor taxonomy.
|
||||
guardrail.output_supervisor.started:
|
||||
label: GRL.001
|
||||
guardrail.result.allow:
|
||||
label: GRL.002
|
||||
guardrail.result.sanitize:
|
||||
label: GRL.003
|
||||
guardrail.result.block:
|
||||
label: GRL.004
|
||||
guardrail.result.retry:
|
||||
label: GRL.005
|
||||
guardrail.result.handover:
|
||||
label: GRL.006
|
||||
guardrail.result.observe:
|
||||
label: GRL.007
|
||||
guardrail.fail_closed:
|
||||
label: GRL.008
|
||||
guardrail.output_supervisor.completed:
|
||||
label: GRL.009
|
||||
|
||||
# Named guardrail events historically emitted as GRL.<RAIL_CODE>.
|
||||
guardrail.input_size: {label: GRL.INPUT_SIZE, aliases: [INPUT_SIZE, SIZE]}
|
||||
guardrail.msk: {label: GRL.MSK, aliases: [MSK, PII]}
|
||||
guardrail.tox: {label: GRL.TOX, aliases: [TOX]}
|
||||
guardrail.pinj: {label: GRL.PINJ, aliases: [PINJ]}
|
||||
guardrail.jailbreak: {label: GRL.JAILBREAK, aliases: [JAILBREAK]}
|
||||
guardrail.vloop: {label: GRL.VLOOP, aliases: [VLOOP, LOOP]}
|
||||
guardrail.dlex_in: {label: GRL.DLEX_IN, aliases: [DLEX_IN]}
|
||||
guardrail.oos: {label: GRL.OOS, aliases: [OOS]}
|
||||
guardrail.coer: {label: GRL.COER, aliases: [COER]}
|
||||
guardrail.msk_out: {label: GRL.MSK_OUT, aliases: [MSK_OUT, OUTPUT_MSK]}
|
||||
guardrail.toxout: {label: GRL.TOXOUT, aliases: [TOXOUT, TOX_OUT]}
|
||||
guardrail.aoferta: {label: GRL.AOFERTA, aliases: [AOFERTA, PROACTIVE_OFFER]}
|
||||
guardrail.dlex_out: {label: GRL.DLEX_OUT, aliases: [DLEX_OUT]}
|
||||
guardrail.aluc_risk: {label: GRL.ALUC_RISK, aliases: [ALUC_RISK, HALLUCINATION_RISK]}
|
||||
guardrail.ret_rel: {label: GRL.RET_REL, aliases: [RET_REL, RETRIEVAL_RELEVANCE]}
|
||||
guardrail.ragsec: {label: GRL.RAGSEC, aliases: [RAGSEC]}
|
||||
guardrail.tool_val: {label: GRL.TOOL_VAL, aliases: [TOOL_VAL, TOOL_VALIDATION]}
|
||||
|
||||
# Historical action-by-name behavior, now declarative.
|
||||
guardrail.revprec:
|
||||
label: GRL.REVPREC
|
||||
action: retry
|
||||
aliases: [REVPREC, PREMATURE_ACTION]
|
||||
guardrail.cmp:
|
||||
label: GRL.CMP
|
||||
action: retry
|
||||
aliases: [CMP, COMPLIANCE]
|
||||
guardrail.sco:
|
||||
label: GRL.SCO
|
||||
action: retry
|
||||
aliases: [SCO]
|
||||
guardrail.gnd:
|
||||
label: GRL.GND
|
||||
action: retry
|
||||
aliases: [GND, GROUNDEDNESS]
|
||||
guardrail.handover:
|
||||
action: handover
|
||||
aliases: [HANDOVER, ATH, HUMAN]
|
||||
|
||||
# Historical FRASEOLOGIA special-case rewrite, now capability-driven.
|
||||
guardrail.fraseologia:
|
||||
label: GRL.FRASEOLOGIA
|
||||
aliases: [FRASEOLOGIA]
|
||||
remediation:
|
||||
type: rewrite
|
||||
max_attempts: 1
|
||||
prompt_id: FALLBACK
|
||||
profile_name: grl
|
||||
component_name: guardrail.fraseologia.rewrite
|
||||
generation_name: guardrail.fraseologia.rewrite
|
||||
@@ -0,0 +1,254 @@
|
||||
from functools import lru_cache
|
||||
from typing import Literal
|
||||
|
||||
from dotenv import load_dotenv
|
||||
from pydantic import Field
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
# Load .env into os.environ as well.
|
||||
# Pydantic Settings reads .env for Settings fields, but parts of the calibrated
|
||||
# guardrails intentionally use os.getenv for compatibility with the original
|
||||
# guardrails package. Loading here keeps both paths consistent.
|
||||
load_dotenv(override=False)
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = SettingsConfigDict(env_file='.env', env_file_encoding='utf-8', extra='ignore')
|
||||
|
||||
APP_NAME: str = 'ai-agent-template'
|
||||
APP_ENV: str = 'local'
|
||||
LOG_LEVEL: str = 'INFO'
|
||||
API_HOST: str = '0.0.0.0'
|
||||
API_PORT: int = 8000
|
||||
CORS_ORIGINS: str = 'http://localhost:5173'
|
||||
|
||||
LLM_PROVIDER: Literal['mock','oci_openai','oci_sdk','openai_compatible'] = 'mock'
|
||||
LLM_TEMPERATURE: float = 0.2
|
||||
LLM_MAX_TOKENS: int = 2048
|
||||
LLM_TIMEOUT_SECONDS: int = 120
|
||||
LLM_PROFILES_PATH: str = './llm_profiles.yaml'
|
||||
# Reasoning controls. When absent from .env, auto is the default.
|
||||
# auto = enable only when the provider/model capability resolver says it is supported.
|
||||
# true = force-enable (the provider still performs SDK/request safety checks).
|
||||
# false = never send reasoning_effort.
|
||||
LLM_REASONING_ENABLED: Literal['auto','true','false'] = 'auto'
|
||||
LLM_REASONING_EFFORT: str | None = None
|
||||
|
||||
OCI_GENAI_BASE_URL: str = ''
|
||||
OCI_GENAI_MODEL: str = 'openai.gpt-4.1'
|
||||
OCI_GENAI_API_KEY: str | None = None
|
||||
OCI_GENAI_PROJECT_OCID: str | None = None
|
||||
# OCI SDK authentication mode.
|
||||
# config_file = ~/.oci/config profile (default/local development)
|
||||
# instance_principal = OCI Instance Principal signer (Compute/OKE without API key)
|
||||
# resource_principal = OCI Resource Principal signer (Functions/resource principal contexts)
|
||||
OCI_AUTH_MODE: Literal['config_file','instance_principal','resource_principal', 'oke_workload_identity'] = 'config_file'
|
||||
OCI_CONFIG_FILE: str = '~/.oci/config'
|
||||
OCI_PROFILE: str = 'DEFAULT'
|
||||
OCI_COMPARTMENT_ID: str | None = None
|
||||
OCI_REGION: str = ''
|
||||
OCI_GENAI_ENDPOINT: str | None = None
|
||||
OCI_EMBEDDING_ENDPOINT: str | None = None
|
||||
|
||||
SESSION_REPOSITORY_PROVIDER: Literal['memory','sqlite','autonomous','oracle','mongodb'] = 'memory'
|
||||
MEMORY_REPOSITORY_PROVIDER: Literal['memory','sqlite','autonomous','oracle','mongodb'] = 'memory'
|
||||
CHECKPOINT_REPOSITORY_PROVIDER: Literal['memory','sqlite','autonomous','oracle','mongodb'] = 'memory'
|
||||
|
||||
# ConversationSummaryMemory: compressão de contexto conversacional.
|
||||
# none = não injeta histórico no prompt
|
||||
# window = injeta somente últimas mensagens
|
||||
# summary = resumo acumulado + últimas mensagens completas
|
||||
ENABLE_CONVERSATION_SUMMARY_MEMORY: bool = False
|
||||
MEMORY_CONTEXT_STRATEGY: Literal['none','window','summary'] = 'window'
|
||||
MEMORY_HISTORY_LIMIT: int = 80
|
||||
MEMORY_RECENT_MESSAGES_LIMIT: int = 8
|
||||
MEMORY_SUMMARY_TRIGGER_MESSAGES: int = 20
|
||||
MEMORY_MAX_SUMMARY_CHARS: int = 6000
|
||||
MEMORY_SUMMARY_USE_LLM: bool = True
|
||||
MEMORY_INJECT_RECENT_MESSAGES: bool = True
|
||||
MEMORY_INJECT_SUMMARY: bool = True
|
||||
|
||||
ENABLE_LONG_TERM_MEMORY: bool = False
|
||||
LONG_TERM_MEMORY_PROVIDER: Literal['memory','sqlite','autonomous','oracle'] = 'sqlite'
|
||||
LONG_TERM_MEMORY_SQLITE_PATH: str | None = None
|
||||
LONG_TERM_MEMORY_TABLE: str = 'agentfw_long_term_memory'
|
||||
LONG_TERM_MEMORY_ORACLE_TABLE: str | None = None
|
||||
LONG_TERM_MEMORY_MAX_CONTEXT_ITEMS: int = 20
|
||||
LONG_TERM_MEMORY_MIN_CONFIDENCE: float = 0.70
|
||||
LONG_TERM_MEMORY_AUTO_EXTRACT: bool = True
|
||||
LONG_TERM_MEMORY_INJECT_CONTEXT: bool = True
|
||||
|
||||
# LangGraph enterprise checkpointing
|
||||
ENABLE_RESILIENT_CHECKPOINTER: bool = True
|
||||
ENABLE_CHECKPOINT_INTEGRITY: bool = True
|
||||
ENABLE_CHECKPOINT_COMPACTION: bool = True
|
||||
CHECKPOINT_COMPACT_EVERY: int = 50
|
||||
CHECKPOINT_KEEP_LAST: int = 20
|
||||
CHECKPOINT_RECOVERY_SCAN_LIMIT: int = 25
|
||||
CHECKPOINT_RETRY_MAX_ATTEMPTS: int = 3
|
||||
CHECKPOINT_RETRY_BASE_DELAY_SECONDS: float = 0.05
|
||||
CHECKPOINT_RETRY_MAX_DELAY_SECONDS: float = 1.0
|
||||
CHECKPOINT_RETRY_JITTER_SECONDS: float = 0.05
|
||||
USAGE_REPOSITORY_PROVIDER: Literal['sqlite','autonomous','oracle'] = 'sqlite'
|
||||
|
||||
ADB_USER: str | None = None
|
||||
ADB_PASSWORD: str | None = None
|
||||
ADB_DSN: str | None = None
|
||||
ADB_WALLET_LOCATION: str | None = None
|
||||
ADB_WALLET_PASSWORD: str | None = None
|
||||
ADB_TABLE_PREFIX: str = 'AGENTFW'
|
||||
|
||||
MONGODB_URI: str = 'mongodb://localhost:27017'
|
||||
MONGODB_DATABASE: str = 'agent_platform'
|
||||
REDIS_URL: str = 'redis://localhost:6379/0'
|
||||
ENABLE_REDIS_CACHE: bool = False
|
||||
CACHE_KEY_PREFIX: str = 'agentfw'
|
||||
|
||||
VECTOR_STORE_PROVIDER: Literal['memory','sqlite','autonomous','oracle','mongodb'] = 'memory'
|
||||
GRAPH_STORE_PROVIDER: Literal['memory','autonomous','oracle'] = 'memory'
|
||||
ORACLE_GRAPH_NAME: str = 'AGENTFW_GRAPH'
|
||||
ORACLE_GRAPH_AUTO_CREATE: bool = False
|
||||
RAG_TOP_K: int = 5
|
||||
SKIP_RAG_WHEN_MCP_SUFFICIENT: bool = True
|
||||
ENABLE_RAG_QUERY_REWRITE: bool = False
|
||||
ENABLE_RAG_CONTEXT_COMPRESSION: bool = False
|
||||
ENABLE_RAG_GENERATION: bool = False
|
||||
EMBEDDING_PROVIDER: Literal['mock','oci'] = 'mock'
|
||||
OCI_EMBEDDING_MODEL: str = 'cohere.embed-multilingual-v3.0'
|
||||
|
||||
ENABLE_LANGFUSE: bool = False
|
||||
LANGFUSE_TRACE_MODE: Literal['verbose','compact'] = 'verbose'
|
||||
LANGFUSE_ROOT_SPAN_NAME: str = 'agent.gateway_message'
|
||||
LANGFUSE_LEGACY_IO_FALLBACK: bool = True
|
||||
LANGFUSE_PUBLIC_KEY: str | None = None
|
||||
LANGFUSE_SECRET_KEY: str | None = None
|
||||
LANGFUSE_HOST: str = 'https://cloud.langfuse.com'
|
||||
MODEL_PRICES_JSON: str | None = None
|
||||
USD_BRL_RATE: str | None = None
|
||||
ENABLE_OTEL: bool = False
|
||||
OTEL_EXPORTER_OTLP_ENDPOINT: str | None = None
|
||||
OTEL_SERVICE_NAME: str = 'ai-agent-template'
|
||||
# Dedicated NOC OpenTelemetry Logs channel. This is separate from trace/span OTel.
|
||||
ENABLE_NOC_OTEL_LOGS: bool = False
|
||||
OTEL_EXPORTER_OTLP_LOGS_ENDPOINT: str | None = None
|
||||
OTEL_EXPORTER_OTLP_HOST_HEADER: str | None = None
|
||||
|
||||
ENABLE_ANALYTICS: bool = False
|
||||
ANALYTICS_PROVIDERS: str = 'oci_streaming'
|
||||
# Framework compatibility registry is loaded by default so legacy agents can
|
||||
# adopt a newer framework without changing their observability/guardrail behavior.
|
||||
OBSERVABILITY_DEFAULT_MAPPING_ENABLED: bool = True
|
||||
OBSERVABILITY_DEFAULT_MAPPING_PATH: str | None = None
|
||||
# Optional agent/deployment overlay applied on top of the framework defaults.
|
||||
OBSERVABILITY_CODE_MAPPING_ENABLED: bool = False
|
||||
OBSERVABILITY_CODE_MAPPING_PATH: str | None = None
|
||||
GCP_PUBSUB_TOPIC_PATH: str | None = None
|
||||
AGENT_PUBSUB_TOPIC: str | None = None
|
||||
GCP_PROJECT_ID: str | None = None
|
||||
GCP_PUBSUB_TOPIC: str | None = None
|
||||
GCP_PUBSUB_TIMEOUT_SECONDS: float = 30.0
|
||||
# Payload shape is a transport concern. Domain-specific adapters must be selected by the embedding application.
|
||||
PUBSUB_PAYLOAD_MODE: Literal['flat','legacy','envelope','wrapped'] = 'flat'
|
||||
# Match the old Observer behavior: NOC.* goes to OTel Logs, not Pub/Sub.
|
||||
PUBSUB_EXCLUDE_NOC: bool = True
|
||||
|
||||
# Automatic Pub/Sub sequence generation.
|
||||
# auto: Redis if configured; otherwise MongoDB if configured; otherwise memory fallback.
|
||||
# mongodb: atomic find_one_and_update/$inc.
|
||||
PUBSUB_SEQUENCE_ENABLED: bool = True
|
||||
PUBSUB_SEQUENCE_PROVIDER: Literal['auto','redis','mongodb','mongo','memory','none'] = 'auto'
|
||||
PUBSUB_SEQUENCE_REDIS_URL: str | None = None
|
||||
PUBSUB_SEQUENCE_MONGODB_URI: str | None = None
|
||||
PUBSUB_SEQUENCE_MONGODB_DATABASE: str | None = None
|
||||
PUBSUB_SEQUENCE_MONGODB_COLLECTION: str = 'observer_sequences'
|
||||
PUBSUB_SEQUENCE_TTL_SECONDS: int = 86400
|
||||
PUBSUB_SEQUENCE_MEMORY_FALLBACK: bool = True
|
||||
PUBSUB_SEQUENCE_KEY_PREFIX: str = 'observer:sequence'
|
||||
|
||||
ANALYTICS_FAIL_SILENT: bool = True
|
||||
|
||||
ENABLE_OCI_STREAMING: bool = False
|
||||
OCI_STREAM_ENDPOINT: str | None = None
|
||||
OCI_STREAM_OCID: str | None = None
|
||||
OCI_STREAM_PARTITION_KEY: str = 'agent-events'
|
||||
|
||||
ENABLE_INPUT_GUARDRAILS: bool = True
|
||||
ENABLE_OUTPUT_GUARDRAILS: bool = True
|
||||
ENABLE_PARALLEL_GUARDRAILS: bool = True
|
||||
GUARDRAILS_FAIL_FAST: bool = True
|
||||
# Optional LLM inference points. Defaults keep the current deterministic behavior.
|
||||
ENABLE_JUDGES: bool = True
|
||||
ENABLE_SUPERVISOR: bool = True
|
||||
ENABLE_OUTPUT_SUPERVISOR: bool = True
|
||||
OUTPUT_SUPERVISOR_MAX_RETRIES: int = 3
|
||||
GUARDRAILS_CONFIG_PATH: str = './config/guardrails.yaml'
|
||||
JUDGES_CONFIG_PATH: str = './config/judges.yaml'
|
||||
PROMPT_POLICY_PATH: str = './config/prompt_policy.yaml'
|
||||
AGENTS_CONFIG_PATH: str = './config/agents.yaml'
|
||||
ROUTING_CONFIG_PATH: str = './config/routing.yaml'
|
||||
ENABLE_LLM_ROUTER: bool = False
|
||||
ROUTING_MODE: Literal['router','supervisor'] = 'router'
|
||||
# Semantic route stickiness. Uses an LLM profile; no regex or language rules.
|
||||
ENABLE_ROUTE_STICKINESS: bool = False
|
||||
ROUTE_STICKINESS_LLM_PROFILE: str = 'route_continuity'
|
||||
ROUTE_STICKINESS_CONFIDENCE_THRESHOLD: float = 0.90
|
||||
ROUTE_STICKINESS_HISTORY_TURNS: int = 2
|
||||
ROUTE_STICKINESS_MAX_TOKENS: int = 80
|
||||
HUMAN_HANDOFF_MESSAGE: str = 'Vou encaminhar seu atendimento para uma pessoa.'
|
||||
END_SESSION_MESSAGE: str = 'Atendimento encerrado. Obrigado pelo contato.'
|
||||
POST_FINALIZE_REPLAY_MESSAGE: str = (
|
||||
'Por aqui finalizamos o tratamento da sua solicitação. '
|
||||
'Aguarde um instante na linha.'
|
||||
)
|
||||
SESSION_ALREADY_ENDED_MESSAGE: str = 'Este atendimento já foi encerrado. Inicie uma nova sessão para continuar.'
|
||||
|
||||
# MCP / Tooling
|
||||
ENABLE_MCP_TOOLS: bool = True
|
||||
ENABLE_MCP_CACHE: bool = True
|
||||
MCP_CACHE_TTL_SECONDS: int = 300
|
||||
MCP_SERVERS_CONFIG_PATH: str = './config/mcp_servers.yaml'
|
||||
TOOLS_CONFIG_PATH: str = './config/tools.yaml'
|
||||
# Opcional. Se ausente, permanecem válidas as políticas legadas de tools.yaml.
|
||||
TOOL_POLICIES_PATH: str | None = './config/tool_policies.yaml'
|
||||
ENABLE_TRANSACTIONAL_WORKFLOWS: bool = False
|
||||
WORKFLOWS_PATH: str = './workflows'
|
||||
IDENTITY_CONFIG_PATH: str = './config/identity.yaml'
|
||||
MCP_PARAMETER_MAPPING_PATH: str = './config/mcp_parameter_mapping.yaml'
|
||||
MCP_TOOL_TIMEOUT_SECONDS: int = 30
|
||||
# When enabled, the framework routes tool calls to the dedicated MCP Gateway
|
||||
# instead of calling individual MCP servers directly. The gateway then owns
|
||||
# server selection, retry, cache and policy enforcement.
|
||||
MCP_GATEWAY_ENABLED: bool = False
|
||||
MCP_GATEWAY_URL: str = 'http://localhost:8300'
|
||||
MCP_GATEWAY_TIMEOUT_SECONDS: int = 60
|
||||
MCP_GATEWAY_TOKEN: str | None = None
|
||||
MCP_GATEWAY_AGENT_ID: str = 'telecom_contas'
|
||||
MCP_GATEWAY_TENANT_ID: str = 'default'
|
||||
|
||||
DEFAULT_CHANNEL: str = 'web'
|
||||
# Agent Framework channel input mode.
|
||||
# embedded = backend may use internal adapters to interpret simple/native payloads.
|
||||
# external = backend accepts only GatewayRequest payloads already normalized by an external Channel Gateway.
|
||||
FRAMEWORK_CHANNEL_INPUT_MODE: Literal['embedded','external'] = 'embedded'
|
||||
# Legacy alias kept for compatibility with older .env files. Prefer FRAMEWORK_CHANNEL_INPUT_MODE.
|
||||
CHANNEL_GATEWAY_MODE: str | None = None
|
||||
ENABLE_VOICE_ADAPTER: bool = True
|
||||
ENABLE_WHATSAPP_ADAPTER: bool = True
|
||||
ENABLE_TEXT_ADAPTER: bool = True
|
||||
|
||||
|
||||
# FIRST-ready runtime options
|
||||
SQLITE_DB_PATH: str = './data/agent_framework.db'
|
||||
ENABLE_SSE: bool = True
|
||||
SSE_KEEPALIVE_SECONDS: float = 15.0
|
||||
SSE_EVENT_REPLAY_LIMIT: int = 100
|
||||
ENABLE_MESSAGE_IDEMPOTENCY: bool = True
|
||||
ENABLE_LOCAL_CACHE: bool = True
|
||||
CACHE_TTL_SECONDS: int = 300
|
||||
CACHE_BACKEND_PROVIDER: Literal['memory','sqlite','autonomous','oracle'] = 'memory'
|
||||
SSE_STORE_PROVIDER: Literal['sqlite','autonomous','oracle'] | None = None
|
||||
|
||||
@lru_cache
|
||||
def get_settings() -> Settings:
|
||||
return Settings()
|
||||
|
||||
settings = get_settings()
|
||||
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,28 @@
|
||||
import json, base64, logging
|
||||
logger=logging.getLogger('agent_framework.streaming')
|
||||
|
||||
class EventPublisher:
|
||||
async def publish(self, event_type: str, payload: dict): ...
|
||||
|
||||
class NoopEventPublisher(EventPublisher):
|
||||
async def publish(self, event_type, payload):
|
||||
logger.info('event.noop %s %s', event_type, payload)
|
||||
|
||||
class OCIStreamingPublisher(EventPublisher):
|
||||
def __init__(self, settings):
|
||||
import oci
|
||||
config = oci.config.from_file(settings.OCI_CONFIG_FILE, settings.OCI_PROFILE)
|
||||
self.client = oci.streaming.StreamClient(config, service_endpoint=settings.OCI_STREAM_ENDPOINT)
|
||||
self.stream_id = settings.OCI_STREAM_OCID
|
||||
self.partition_key = settings.OCI_STREAM_PARTITION_KEY
|
||||
async def publish(self, event_type, payload):
|
||||
import oci
|
||||
body = json.dumps({'type': event_type, 'payload': payload}, default=str).encode()
|
||||
entry = oci.streaming.models.PutMessagesDetailsEntry(key=self.partition_key.encode(), value=body)
|
||||
details = oci.streaming.models.PutMessagesDetails(messages=[entry])
|
||||
self.client.put_messages(self.stream_id, details)
|
||||
|
||||
def create_event_publisher(settings):
|
||||
if settings.ENABLE_OCI_STREAMING and settings.OCI_STREAM_ENDPOINT and settings.OCI_STREAM_OCID:
|
||||
return OCIStreamingPublisher(settings)
|
||||
return NoopEventPublisher()
|
||||
@@ -0,0 +1,47 @@
|
||||
from __future__ import annotations
|
||||
|
||||
"""Extension SPI for agent-owned guardrails and judges.
|
||||
|
||||
The framework owns execution, telemetry and lifecycle. Agents may contribute
|
||||
classes through YAML using ``type: external`` and ``class: module:Class``.
|
||||
No agent/domain package is imported unless explicitly declared in configuration.
|
||||
"""
|
||||
|
||||
from importlib import import_module
|
||||
from typing import Any
|
||||
|
||||
|
||||
def load_external_class(path: str) -> type[Any]:
|
||||
value = str(path or "").strip()
|
||||
if not value:
|
||||
raise ValueError("External component requires 'class: module:ClassName'")
|
||||
if ':' in value:
|
||||
module_name, class_name = value.rsplit(':', 1)
|
||||
elif '.' in value:
|
||||
module_name, class_name = value.rsplit('.', 1)
|
||||
else:
|
||||
raise ValueError(f"Invalid external class path: {value}")
|
||||
module = import_module(module_name)
|
||||
cls = getattr(module, class_name, None)
|
||||
if cls is None or not isinstance(cls, type):
|
||||
raise ValueError(f"External class not found: {value}")
|
||||
return cls
|
||||
|
||||
|
||||
def instantiate_external(path: str, *, kwargs: dict[str, Any] | None = None, injected: dict[str, Any] | None = None) -> Any:
|
||||
cls = load_external_class(path)
|
||||
params = dict(kwargs or {})
|
||||
for key, value in (injected or {}).items():
|
||||
params.setdefault(key, value)
|
||||
try:
|
||||
return cls(**params)
|
||||
except TypeError:
|
||||
# Backward-friendly path for simple plugins with no constructor args.
|
||||
if params:
|
||||
obj = cls()
|
||||
for key, value in params.items():
|
||||
if not hasattr(obj, key):
|
||||
continue
|
||||
setattr(obj, key, value)
|
||||
return obj
|
||||
raise
|
||||
@@ -0,0 +1,27 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
|
||||
def get_gateway_model_policy(state: dict[str, Any]) -> dict[str, Any] | None:
|
||||
metadata = state.get("metadata") or {}
|
||||
policy = metadata.get("model_policy")
|
||||
return policy if isinstance(policy, dict) else None
|
||||
|
||||
|
||||
def apply_gateway_model_policy_to_llm_kwargs(
|
||||
state: dict[str, Any],
|
||||
fallback_profile: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
policy = get_gateway_model_policy(state)
|
||||
if not policy:
|
||||
return fallback_profile or {}
|
||||
|
||||
params = dict(policy.get("parameters") or {})
|
||||
if policy.get("model"):
|
||||
params["model"] = policy["model"]
|
||||
if policy.get("provider"):
|
||||
params["provider"] = policy["provider"]
|
||||
if policy.get("profile"):
|
||||
params["profile"] = policy["profile"]
|
||||
return params
|
||||
@@ -0,0 +1,3 @@
|
||||
from .mcp_gateway_client import MCPGatewayClient
|
||||
|
||||
__all__ = ["MCPGatewayClient"]
|
||||
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,50 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
class MCPGatewayClient:
|
||||
def __init__(self, base_url: str, token: str | None = None, timeout_seconds: int = 60):
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.token = token
|
||||
self.timeout_seconds = timeout_seconds
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
return {"Authorization": f"Bearer {self.token}"} if self.token else {}
|
||||
|
||||
async def list_tools(self) -> dict[str, Any]:
|
||||
async with httpx.AsyncClient(timeout=self.timeout_seconds) as client:
|
||||
response = await client.get(f"{self.base_url}/v1/tools", headers=self._headers())
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def invoke_tool(
|
||||
self,
|
||||
*,
|
||||
tenant_id: str,
|
||||
agent_id: str,
|
||||
channel: str | None,
|
||||
tool_name: str,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
business_context: dict[str, Any] | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
payload = {
|
||||
"tenant_id": tenant_id,
|
||||
"agent_id": agent_id,
|
||||
"channel": channel,
|
||||
"tool_name": tool_name,
|
||||
"arguments": arguments or {},
|
||||
"business_context": business_context or {},
|
||||
"metadata": metadata or {},
|
||||
}
|
||||
async with httpx.AsyncClient(timeout=self.timeout_seconds) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/v1/tools/{tool_name}/invoke",
|
||||
json=payload,
|
||||
headers=self._headers(),
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
@@ -0,0 +1,25 @@
|
||||
from .client import BackendClient
|
||||
from .config import BackendRegistry
|
||||
from .models import (
|
||||
BackendCallResult,
|
||||
BackendDefinition,
|
||||
BackendRegistryConfig,
|
||||
GlobalRouteDecision,
|
||||
GlobalRouteRequest,
|
||||
GlobalSessionState,
|
||||
)
|
||||
from .router import GlobalSupervisorRouter
|
||||
from .session_store import InMemoryGlobalSessionStore
|
||||
|
||||
__all__ = [
|
||||
"BackendClient",
|
||||
"BackendRegistry",
|
||||
"BackendCallResult",
|
||||
"BackendDefinition",
|
||||
"BackendRegistryConfig",
|
||||
"GlobalRouteDecision",
|
||||
"GlobalRouteRequest",
|
||||
"GlobalSessionState",
|
||||
"GlobalSupervisorRouter",
|
||||
"InMemoryGlobalSessionStore",
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,60 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from .models import BackendCallResult, BackendDefinition, GlobalRouteDecision
|
||||
|
||||
|
||||
class BackendClient:
|
||||
def __init__(self, timeout_seconds: float = 120.0):
|
||||
self.timeout_seconds = timeout_seconds
|
||||
|
||||
async def call_message(
|
||||
self,
|
||||
backend: BackendDefinition,
|
||||
request_payload: dict[str, Any],
|
||||
route_decision: GlobalRouteDecision,
|
||||
use_sse: bool = False,
|
||||
) -> BackendCallResult:
|
||||
path = backend.sse_message_path if use_sse else backend.message_path
|
||||
url = f"{backend.base_url}{path}"
|
||||
payload = dict(request_payload)
|
||||
# Mantém compatibilidade com agent_template_backend.
|
||||
payload.setdefault("agent_id", backend.default_agent_id)
|
||||
payload.setdefault("tenant_id", request_payload.get("tenant_id"))
|
||||
inner = payload.setdefault("payload", {}) if isinstance(payload.get("payload"), dict) else None
|
||||
if inner is not None:
|
||||
inner.setdefault("selected_backend", backend.backend_id)
|
||||
inner.setdefault("global_route_decision", route_decision.model_dump(mode="json"))
|
||||
started = time.time()
|
||||
async with httpx.AsyncClient(timeout=self.timeout_seconds) as client:
|
||||
resp = await client.post(url, json=payload)
|
||||
elapsed_ms = int((time.time() - started) * 1000)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return BackendCallResult(
|
||||
backend_id=backend.backend_id,
|
||||
backend_url=backend.base_url,
|
||||
status_code=resp.status_code,
|
||||
response=data,
|
||||
route_decision=route_decision,
|
||||
elapsed_ms=elapsed_ms,
|
||||
)
|
||||
|
||||
async def health(self, backend: BackendDefinition) -> dict[str, Any]:
|
||||
url = f"{backend.base_url}{backend.health_path}"
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
try:
|
||||
resp = await client.get(url)
|
||||
return {"backend_id": backend.backend_id, "status_code": resp.status_code, "ok": resp.is_success, "body": self._safe_json(resp)}
|
||||
except Exception as exc:
|
||||
return {"backend_id": backend.backend_id, "ok": False, "error": str(exc)}
|
||||
|
||||
def _safe_json(self, resp: httpx.Response) -> Any:
|
||||
try:
|
||||
return resp.json()
|
||||
except Exception:
|
||||
return resp.text[:500]
|
||||
@@ -0,0 +1,65 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import yaml
|
||||
|
||||
from .models import BackendDefinition, BackendRegistryConfig
|
||||
|
||||
|
||||
class BackendRegistry:
|
||||
def __init__(self, config: BackendRegistryConfig):
|
||||
self.config = config
|
||||
self.backends: dict[str, BackendDefinition] = {
|
||||
b.backend_id: b for b in config.backends if b.enabled
|
||||
}
|
||||
if not self.backends:
|
||||
raise ValueError("Nenhum backend habilitado no registry do Global Supervisor.")
|
||||
|
||||
@classmethod
|
||||
def from_yaml(cls, path: str | Path) -> "BackendRegistry":
|
||||
p = Path(path)
|
||||
data = yaml.safe_load(p.read_text(encoding="utf-8")) or {}
|
||||
raw_backends = data.get("backends") or []
|
||||
# Aceita lista ou dict para facilitar edição humana do YAML.
|
||||
if isinstance(raw_backends, dict):
|
||||
normalized = []
|
||||
for backend_id, value in raw_backends.items():
|
||||
item = dict(value or {})
|
||||
item.setdefault("backend_id", backend_id)
|
||||
normalized.append(item)
|
||||
raw_backends = normalized
|
||||
config = BackendRegistryConfig(
|
||||
default_backend=data.get("default_backend"),
|
||||
backends=[BackendDefinition(**b) for b in raw_backends],
|
||||
)
|
||||
return cls(config)
|
||||
|
||||
def get(self, backend_id: str) -> BackendDefinition:
|
||||
try:
|
||||
return self.backends[backend_id]
|
||||
except KeyError as exc:
|
||||
raise KeyError(f"Backend não registrado ou desabilitado: {backend_id}") from exc
|
||||
|
||||
def default(self) -> BackendDefinition:
|
||||
if self.config.default_backend and self.config.default_backend in self.backends:
|
||||
return self.backends[self.config.default_backend]
|
||||
return sorted(self.backends.values(), key=lambda b: b.priority)[0]
|
||||
|
||||
def list(self) -> list[BackendDefinition]:
|
||||
return sorted(self.backends.values(), key=lambda b: (b.priority, b.backend_id))
|
||||
|
||||
def describe_for_prompt(self) -> str:
|
||||
lines: list[str] = []
|
||||
for b in self.list():
|
||||
lines.append(
|
||||
f"- {b.backend_id}: {b.description} | domínios={', '.join(b.domains)} | exemplos={'; '.join(b.examples[:3])}"
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
def as_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"default_backend": self.config.default_backend,
|
||||
"backends": [b.model_dump(mode="json") for b in self.list()],
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
RoutingMode = Literal["router", "supervisor", "hybrid"]
|
||||
|
||||
|
||||
class BackendDefinition(BaseModel):
|
||||
"""Contrato de um backend de agente registrado no Global Supervisor."""
|
||||
|
||||
backend_id: str = Field(..., description="Identificador lógico. Ex.: contas, ofertas, suporte")
|
||||
name: str | None = None
|
||||
url: str = Field(..., description="Base URL do backend, sem barra final")
|
||||
description: str = ""
|
||||
domains: list[str] = Field(default_factory=list)
|
||||
keywords: list[str] = Field(default_factory=list)
|
||||
examples: list[str] = Field(default_factory=list)
|
||||
priority: int = 100
|
||||
enabled: bool = True
|
||||
health_path: str = "/health"
|
||||
message_path: str = "/gateway/message"
|
||||
sse_message_path: str = "/gateway/message/sse"
|
||||
events_path_template: str = "/gateway/events/{session_id}"
|
||||
default_agent_id: str | None = None
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def base_url(self) -> str:
|
||||
return self.url.rstrip("/")
|
||||
|
||||
|
||||
class BackendRegistryConfig(BaseModel):
|
||||
default_backend: str | None = None
|
||||
backends: list[BackendDefinition] = Field(default_factory=list)
|
||||
|
||||
|
||||
class GlobalRouteRequest(BaseModel):
|
||||
channel: str = "web"
|
||||
payload: dict[str, Any] = Field(default_factory=dict)
|
||||
tenant_id: str | None = None
|
||||
session_id: str | None = None
|
||||
current_backend: str | None = None
|
||||
force_backend: str | None = None
|
||||
mode: RoutingMode | None = None
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class GlobalRouteDecision(BaseModel):
|
||||
backend_id: str
|
||||
confidence: float = 0.0
|
||||
reason: str = ""
|
||||
mode: RoutingMode = "hybrid"
|
||||
used_llm: bool = False
|
||||
keep_active_backend: bool = False
|
||||
candidates: list[dict[str, Any]] = Field(default_factory=list)
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class BackendCallResult(BaseModel):
|
||||
backend_id: str
|
||||
backend_url: str
|
||||
status_code: int
|
||||
response: dict[str, Any]
|
||||
route_decision: GlobalRouteDecision
|
||||
elapsed_ms: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlobalSessionState:
|
||||
session_id: str
|
||||
tenant_id: str = "default"
|
||||
active_backend: str | None = None
|
||||
active_domain: str | None = None
|
||||
turn_count: int = 0
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
@@ -0,0 +1,258 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from .config import BackendRegistry
|
||||
from .models import BackendDefinition, GlobalRouteDecision, GlobalRouteRequest, RoutingMode
|
||||
from .session_store import InMemoryGlobalSessionStore
|
||||
|
||||
logger = logging.getLogger("agent_framework.global_supervisor")
|
||||
|
||||
_TERMINAL_WORDS = {
|
||||
"obrigado", "obrigada", "valeu", "tchau", "encerrar", "fim", "cancelar atendimento"
|
||||
}
|
||||
|
||||
|
||||
class GlobalSupervisorRouter:
|
||||
"""Roteador global entre backends.
|
||||
|
||||
Modos:
|
||||
- router: usa regras/keywords/domínios do YAML.
|
||||
- supervisor: usa LLM para escolher backend.
|
||||
- hybrid: mantém backend ativo quando coerente; usa router; chama LLM quando ambíguo.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
registry: BackendRegistry,
|
||||
llm: Any | None = None,
|
||||
session_store: InMemoryGlobalSessionStore | None = None,
|
||||
mode: RoutingMode = "hybrid",
|
||||
keep_active_backend: bool = True,
|
||||
use_supervisor_on_conflict: bool = True,
|
||||
min_router_confidence: float = 0.55,
|
||||
):
|
||||
self.registry = registry
|
||||
self.llm = llm
|
||||
self.session_store = session_store or InMemoryGlobalSessionStore()
|
||||
self.mode = mode
|
||||
self.keep_active_backend = keep_active_backend
|
||||
self.use_supervisor_on_conflict = use_supervisor_on_conflict
|
||||
self.min_router_confidence = min_router_confidence
|
||||
|
||||
async def route(self, request: GlobalRouteRequest) -> GlobalRouteDecision:
|
||||
mode = request.mode or self.mode
|
||||
session_id = self._session_id(request)
|
||||
tenant_id = request.tenant_id or request.payload.get("tenant_id") or "default"
|
||||
|
||||
if request.force_backend:
|
||||
decision = self._forced_decision(request.force_backend, mode)
|
||||
await self.session_store.set_active_backend(session_id, decision.backend_id, tenant_id, forced=True)
|
||||
return decision
|
||||
|
||||
state = await self.session_store.get(session_id)
|
||||
text = self._extract_text(request).strip()
|
||||
|
||||
if mode == "router":
|
||||
decision = self._route_by_rules(text, mode)
|
||||
elif mode == "supervisor":
|
||||
decision = await self._route_by_llm(text, request, mode)
|
||||
else:
|
||||
decision = await self._route_hybrid(text, request, state, mode)
|
||||
|
||||
await self.session_store.set_active_backend(
|
||||
session_id,
|
||||
decision.backend_id,
|
||||
tenant_id,
|
||||
last_reason=decision.reason,
|
||||
last_mode=decision.mode,
|
||||
last_confidence=decision.confidence,
|
||||
)
|
||||
return decision
|
||||
|
||||
async def _route_hybrid(self, text: str, request: GlobalRouteRequest, state, mode: RoutingMode) -> GlobalRouteDecision:
|
||||
# Se a conversa já tem backend ativo e a mensagem parece continuação curta, mantenha.
|
||||
active_backend = request.current_backend or (state.active_backend if state else None)
|
||||
if self.keep_active_backend and active_backend and active_backend in self.registry.backends:
|
||||
if self._looks_like_followup(text):
|
||||
return GlobalRouteDecision(
|
||||
backend_id=active_backend,
|
||||
confidence=0.78,
|
||||
reason="Mensagem parece continuação; mantendo backend ativo da sessão.",
|
||||
mode=mode,
|
||||
keep_active_backend=True,
|
||||
)
|
||||
|
||||
rule_decision = self._route_by_rules(text, mode)
|
||||
if rule_decision.confidence >= self.min_router_confidence:
|
||||
return rule_decision
|
||||
|
||||
if self.use_supervisor_on_conflict and self.llm:
|
||||
llm_decision = await self._route_by_llm(text, request, mode, fallback=rule_decision)
|
||||
return llm_decision
|
||||
|
||||
if active_backend and active_backend in self.registry.backends:
|
||||
return GlobalRouteDecision(
|
||||
backend_id=active_backend,
|
||||
confidence=0.50,
|
||||
reason="Router ficou ambíguo; mantendo backend ativo por política híbrida.",
|
||||
mode=mode,
|
||||
keep_active_backend=True,
|
||||
candidates=rule_decision.candidates,
|
||||
)
|
||||
return rule_decision
|
||||
|
||||
def _route_by_rules(self, text: str, mode: RoutingMode) -> GlobalRouteDecision:
|
||||
normalized = self._normalize(text)
|
||||
scored: list[tuple[float, BackendDefinition, list[str]]] = []
|
||||
for backend in self.registry.list():
|
||||
hits: list[str] = []
|
||||
score = 0.0
|
||||
for kw in backend.keywords:
|
||||
nkw = self._normalize(kw)
|
||||
if nkw and nkw in normalized:
|
||||
hits.append(kw)
|
||||
score += 1.0
|
||||
for domain in backend.domains:
|
||||
nd = self._normalize(domain)
|
||||
if nd and nd in normalized:
|
||||
hits.append(domain)
|
||||
score += 0.7
|
||||
if score:
|
||||
# prioridade menor aumenta levemente confiança
|
||||
score += max(0, (200 - backend.priority)) / 1000
|
||||
scored.append((score, backend, hits))
|
||||
|
||||
scored.sort(key=lambda x: (-x[0], x[1].priority, x[1].backend_id))
|
||||
best_score, best_backend, hits = scored[0] if scored else (0.0, self.registry.default(), [])
|
||||
if best_score <= 0:
|
||||
best_backend = self.registry.default()
|
||||
confidence = 0.25
|
||||
reason = "Nenhuma regra forte encontrada; usando backend default."
|
||||
else:
|
||||
# normalização simples para 0..1
|
||||
confidence = min(0.95, 0.35 + best_score / 4)
|
||||
reason = f"Backend escolhido por regras: matches={hits}."
|
||||
candidates = [
|
||||
{"backend_id": b.backend_id, "score": round(s, 3), "matches": h}
|
||||
for s, b, h in scored[:5]
|
||||
]
|
||||
return GlobalRouteDecision(
|
||||
backend_id=best_backend.backend_id,
|
||||
confidence=confidence,
|
||||
reason=reason,
|
||||
mode=mode,
|
||||
used_llm=False,
|
||||
candidates=candidates,
|
||||
)
|
||||
|
||||
async def _route_by_llm(
|
||||
self,
|
||||
text: str,
|
||||
request: GlobalRouteRequest,
|
||||
mode: RoutingMode,
|
||||
fallback: GlobalRouteDecision | None = None,
|
||||
) -> GlobalRouteDecision:
|
||||
if not self.llm:
|
||||
return fallback or self._route_by_rules(text, mode)
|
||||
prompt = self._build_supervisor_prompt(text, request)
|
||||
try:
|
||||
raw = await self.llm.ainvoke([
|
||||
{"role": "system", "content": "Você é um supervisor global de backends. Responda somente JSON válido."},
|
||||
{"role": "user", "content": prompt},
|
||||
], temperature=0, profile_name="supervisor", component_name="supervisor", generation_name="llm.supervisor")
|
||||
data = self._parse_json(raw)
|
||||
backend_id = str(data.get("backend") or data.get("backend_id") or "").strip()
|
||||
if backend_id not in self.registry.backends:
|
||||
raise ValueError(f"LLM retornou backend inválido: {backend_id!r}")
|
||||
return GlobalRouteDecision(
|
||||
backend_id=backend_id,
|
||||
confidence=float(data.get("confidence", 0.75)),
|
||||
reason=str(data.get("reason", "Selecionado pelo supervisor LLM.")),
|
||||
mode=mode,
|
||||
used_llm=True,
|
||||
candidates=(fallback.candidates if fallback else []),
|
||||
metadata={"raw_llm": raw},
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception("Falha no supervisor LLM; usando fallback/router: %s", exc)
|
||||
decision = fallback or self._route_by_rules(text, mode)
|
||||
decision.reason = f"Fallback após falha do supervisor LLM: {decision.reason}"
|
||||
return decision
|
||||
|
||||
def _build_supervisor_prompt(self, text: str, request: GlobalRouteRequest) -> str:
|
||||
history = request.payload.get("history") or request.metadata.get("history") or []
|
||||
return (
|
||||
"Escolha o backend mais adequado para atender a mensagem do usuário.\n\n"
|
||||
"Backends disponíveis:\n"
|
||||
f"{self.registry.describe_for_prompt()}\n\n"
|
||||
"Mensagem atual:\n"
|
||||
f"{text}\n\n"
|
||||
"Histórico/metadata resumidos:\n"
|
||||
f"{json.dumps({'history': history[-6:] if isinstance(history, list) else history, 'metadata': request.metadata}, ensure_ascii=False)[:4000]}\n\n"
|
||||
"Retorne somente JSON neste formato:\n"
|
||||
'{"backend":"<id>","confidence":0.0,"reason":"..."}'
|
||||
)
|
||||
|
||||
def _forced_decision(self, backend_id: str, mode: RoutingMode) -> GlobalRouteDecision:
|
||||
self.registry.get(backend_id)
|
||||
return GlobalRouteDecision(
|
||||
backend_id=backend_id,
|
||||
confidence=1.0,
|
||||
reason="Backend forçado na requisição.",
|
||||
mode=mode,
|
||||
used_llm=False,
|
||||
)
|
||||
|
||||
def _looks_like_followup(self, text: str) -> bool:
|
||||
n = self._normalize(text)
|
||||
if not n:
|
||||
return True
|
||||
if n in _TERMINAL_WORDS:
|
||||
return False
|
||||
tokens = n.split()
|
||||
followup_markers = ["esse", "essa", "isso", "valor", "ele", "ela", "tambem", "e ", "entao", "nesse", "nessa"]
|
||||
return len(tokens) <= 6 or any(marker in n for marker in followup_markers)
|
||||
|
||||
def _extract_text(self, request: GlobalRouteRequest) -> str:
|
||||
payload = request.payload or {}
|
||||
for key in ("text", "message", "input", "user_text"):
|
||||
if payload.get(key):
|
||||
return str(payload[key])
|
||||
if isinstance(payload.get("payload"), dict):
|
||||
inner = payload["payload"]
|
||||
for key in ("text", "message", "input", "user_text"):
|
||||
if inner.get(key):
|
||||
return str(inner[key])
|
||||
return str(payload)
|
||||
|
||||
def _session_id(self, request: GlobalRouteRequest) -> str:
|
||||
payload = request.payload or {}
|
||||
return (
|
||||
request.session_id
|
||||
or payload.get("session_id")
|
||||
or payload.get("conversation_key")
|
||||
or request.metadata.get("session_id")
|
||||
or "global-default-session"
|
||||
)
|
||||
|
||||
def _normalize(self, text: str) -> str:
|
||||
text = text.lower()
|
||||
text = re.sub(r"[^a-z0-9áàâãéêíóôõúçñ\s]", " ", text)
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
return text.strip()
|
||||
|
||||
def _parse_json(self, raw: Any) -> dict[str, Any]:
|
||||
if isinstance(raw, dict):
|
||||
return raw
|
||||
text = str(raw).strip()
|
||||
if text.startswith("```"):
|
||||
text = re.sub(r"^```(?:json)?", "", text).strip()
|
||||
text = re.sub(r"```$", "", text).strip()
|
||||
match = re.search(r"\{.*\}", text, flags=re.S)
|
||||
if match:
|
||||
text = match.group(0)
|
||||
return json.loads(text)
|
||||
@@ -0,0 +1,61 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from dataclasses import asdict
|
||||
|
||||
from .models import GlobalSessionState
|
||||
|
||||
|
||||
class InMemoryGlobalSessionStore:
|
||||
"""Store simples para o Agent Gateway.
|
||||
|
||||
Em produção, use o mesmo repositório compartilhado dos backends
|
||||
(Autonomous DB/Mongo/Redis) para manter handoff entre serviços.
|
||||
"""
|
||||
|
||||
def __init__(self, ttl_seconds: int = 3600):
|
||||
self.ttl_seconds = ttl_seconds
|
||||
self._data: dict[str, tuple[float, GlobalSessionState]] = {}
|
||||
|
||||
async def get(self, session_id: str) -> GlobalSessionState | None:
|
||||
item = self._data.get(session_id)
|
||||
if not item:
|
||||
return None
|
||||
ts, state = item
|
||||
if time.time() - ts > self.ttl_seconds:
|
||||
self._data.pop(session_id, None)
|
||||
return None
|
||||
return state
|
||||
|
||||
async def upsert(self, state: GlobalSessionState) -> None:
|
||||
state.turn_count += 1
|
||||
self._data[state.session_id] = (time.time(), state)
|
||||
|
||||
async def set_active_backend(self, session_id: str, backend_id: str, tenant_id: str = "default", **metadata) -> GlobalSessionState:
|
||||
state = await self.get(session_id) or GlobalSessionState(session_id=session_id, tenant_id=tenant_id)
|
||||
state.active_backend = backend_id
|
||||
state.metadata.update(metadata)
|
||||
await self.upsert(state)
|
||||
return state
|
||||
|
||||
async def dump(self) -> dict:
|
||||
return {k: asdict(v[1]) for k, v in self._data.items()}
|
||||
|
||||
async def rename_session(
|
||||
self,
|
||||
old_session_id: str,
|
||||
new_session_id: str
|
||||
) -> GlobalSessionState | None:
|
||||
|
||||
item = self._data.pop(old_session_id, None)
|
||||
|
||||
if not item:
|
||||
return None
|
||||
|
||||
ts, state = item
|
||||
|
||||
state.session_id = new_session_id
|
||||
|
||||
self._data[new_session_id] = (ts, state)
|
||||
|
||||
return state
|
||||
@@ -0,0 +1,60 @@
|
||||
from .base import Guardrail, RailDecision
|
||||
from .pipeline import GuardrailPipeline
|
||||
from .llm_rails import LLMGuardrailRail, LLMOutputGRLRail
|
||||
from .rails import (
|
||||
ComplianceRail,
|
||||
DataLeakageInputRail,
|
||||
DataLeakageOutputRail,
|
||||
GroundednessRail,
|
||||
HallucinationRiskRail,
|
||||
JailbreakRail,
|
||||
LoopRail,
|
||||
MessageSizeRail,
|
||||
OutOfScopeRail,
|
||||
OutputPiiMaskRail,
|
||||
OutputToxicitySanitizationRail,
|
||||
PiiMaskRail,
|
||||
PrematureActionRail,
|
||||
ProactiveOfferRail,
|
||||
PromptInjectionRail,
|
||||
RagSecurityRail,
|
||||
RetrievalRelevanceRail,
|
||||
ToolValidationRail,
|
||||
ToxicityRail,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Guardrail",
|
||||
"RailDecision",
|
||||
"GuardrailPipeline",
|
||||
"LLMGuardrailRail",
|
||||
"LLMOutputGRLRail",
|
||||
"PiiMaskRail",
|
||||
"OutputPiiMaskRail",
|
||||
"OutputToxicitySanitizationRail",
|
||||
"ToxicityRail",
|
||||
"PromptInjectionRail",
|
||||
"JailbreakRail",
|
||||
"MessageSizeRail",
|
||||
"OutOfScopeRail",
|
||||
"LoopRail",
|
||||
"PrematureActionRail",
|
||||
"ProactiveOfferRail",
|
||||
"RagSecurityRail",
|
||||
"ComplianceRail",
|
||||
"DataLeakageInputRail",
|
||||
"DataLeakageOutputRail",
|
||||
"GroundednessRail",
|
||||
"HallucinationRiskRail",
|
||||
"RetrievalRelevanceRail",
|
||||
"ToolValidationRail",
|
||||
"ParallelRailExecutor",
|
||||
"ParallelRailExecution",
|
||||
]
|
||||
from .rail_action import RailAction
|
||||
from .rail_result import RailResult
|
||||
from .rail_decision import RailDecisionV2
|
||||
from .output_supervisor import OutputSupervisor
|
||||
from .custom_rails import CustomRails
|
||||
|
||||
from .parallel_executor import ParallelRailExecutor, ParallelRailExecution
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user