diff --git a/core/src/agentforge_core/runtime/orchestrator.py b/core/src/agentforge_core/runtime/orchestrator.py new file mode 100644 index 0000000..6e9e817 --- /dev/null +++ b/core/src/agentforge_core/runtime/orchestrator.py @@ -0,0 +1,164 @@ +"""Orchestrator: envuelve build_graph + ainvoke/resume + serialización a ``AgentExecution``. + +Es el único punto de entrada al runtime: el router de FastAPI lo usa para lanzar +ejecuciones y reanudar pausas Human-in-the-Loop. Cada llamada abre su propio +``AsyncSqliteSaver`` (context manager) sobre ``data_dir/checkpoints.sqlite``, así que +el estado de un ``awaiting_approval`` sobrevive a un reinicio del proceso: basta crear +otro ``AgentOrchestrator`` apuntando al mismo ``data_dir`` y llamar a ``resume``. +""" + +from __future__ import annotations + +from datetime import UTC, datetime +from pathlib import Path +from typing import Any +from uuid import UUID, uuid4 + +import structlog +from langgraph.types import Command + +from agentforge_core.domain.agent import AgentDefinition +from agentforge_core.domain.execution import AgentExecution, DecisionStep, ProposedAction +from agentforge_core.domain.guardrail import GuardrailViolation +from agentforge_core.domain.policy import PolicyDefinition +from agentforge_core.guardrails.base import GuardrailEngine +from agentforge_core.llm.base import LLMProvider +from agentforge_core.runtime.checkpointer import build_checkpointer +from agentforge_core.runtime.graph import build_graph + +log = structlog.get_logger(__name__) + +_TERMINAL_STATUSES = frozenset({"completed", "failed", "blocked_by_guardrail"}) + + +def _thread_config(trace_id: UUID) -> dict[str, Any]: + return {"configurable": {"thread_id": str(trace_id)}} + + +class AgentOrchestrator: + """Punto único de entrada para invocar agentes y reanudar pausas HITL.""" + + def __init__( + self, *, provider: LLMProvider, engine: GuardrailEngine, data_dir: Path + ) -> None: + self._provider = provider + self._engine = engine + self._data_dir = data_dir + + async def invoke( + self, + *, + agent_def: AgentDefinition, + policy: PolicyDefinition, + user_input: str, + trace_id: UUID | None = None, + ) -> AgentExecution: + """Lanza una ejecución. Si hay acciones que requieren aprobación, devuelve + ``status="awaiting_approval"`` y el grafo queda pausado en el checkpointer.""" + tid = trace_id or uuid4() + started_at = datetime.now(UTC) + initial: dict[str, Any] = { + "trace_id": str(tid), + "agent_name": agent_def.name, + "agent_version": agent_def.version, + "user_input": user_input, + "messages": [], + "raw_llm_output": None, + "parsed_output": None, + "proposed_actions": [], + "violations": [], + "decision_path": [], + "status": "running", + "error": None, + "human_decision": None, + "final_output": None, + } + async with build_checkpointer(self._data_dir) as checkpointer: + graph = build_graph( + agent_def=agent_def, + policy=policy, + provider=self._provider, + engine=self._engine, + checkpointer=checkpointer, + ) + crashed = False + try: + await graph.ainvoke(initial, config=_thread_config(tid)) + except Exception: + log.exception("agent_invoke_failed", trace_id=str(tid)) + crashed = True + return await self._snapshot(graph, agent_def, tid, started_at, crashed=crashed) + + async def resume( + self, + *, + agent_def: AgentDefinition, + policy: PolicyDefinition, + trace_id: UUID, + decision: dict[str, Any], + ) -> AgentExecution: + """Reanuda una ejecución pausada en HITL con la decisión del operador.""" + started_at = datetime.now(UTC) + async with build_checkpointer(self._data_dir) as checkpointer: + graph = build_graph( + agent_def=agent_def, + policy=policy, + provider=self._provider, + engine=self._engine, + checkpointer=checkpointer, + ) + crashed = False + try: + await graph.ainvoke(Command(resume=decision), config=_thread_config(trace_id)) + except Exception: + log.exception("agent_resume_failed", trace_id=str(trace_id)) + crashed = True + return await self._snapshot(graph, agent_def, trace_id, started_at, crashed=crashed) + + async def _snapshot( + self, + graph: Any, + agent_def: AgentDefinition, + trace_id: UUID, + started_at: datetime, + *, + crashed: bool, + ) -> AgentExecution: + """Lee el estado del checkpointer y lo serializa a ``AgentExecution``.""" + state = await graph.aget_state(_thread_config(trace_id)) + values: dict[str, Any] = state.values or {} + status: str = values.get("status", "running") + if crashed and status not in _TERMINAL_STATUSES: + status = "failed" + elif state.next and status not in _TERMINAL_STATUSES: + # LangGraph reporta nodos pendientes → pausado en interrupt() (approve_gate). + status = "awaiting_approval" + + proposed = [ProposedAction.model_validate(a) for a in values.get("proposed_actions", [])] + violations = [GuardrailViolation.model_validate(v) for v in values.get("violations", [])] + decision_path = [DecisionStep.model_validate(s) for s in values.get("decision_path", [])] + needs_human = ( + [ + a + for a in proposed + if a.risk_score >= agent_def.risk_threshold_for_hitl or a.requires_approval + ] + if status == "awaiting_approval" + else None + ) + error = values.get("error") or ("internal_error" if crashed else None) + finished_at = datetime.now(UTC) if status in _TERMINAL_STATUSES else None + return AgentExecution( + trace_id=trace_id, + agent_name=agent_def.name, + agent_version=agent_def.version, + status=status, + started_at=started_at, + finished_at=finished_at, + decision_path=decision_path, + violations=violations, + proposed_actions=proposed, + needs_human_for=needs_human, + final_output=values.get("final_output"), + error=error, + ) diff --git a/tests/unit/test_runtime_orchestrator.py b/tests/unit/test_runtime_orchestrator.py new file mode 100644 index 0000000..6c28452 --- /dev/null +++ b/tests/unit/test_runtime_orchestrator.py @@ -0,0 +1,138 @@ +"""Tests del AgentOrchestrator: invoke, pausa HITL, resume (aprobar/rechazar) y persistencia.""" + +from datetime import UTC, datetime +from pathlib import Path + +import pytest + +from agentforge_core.domain.agent import AgentDefinition, LLMConfig +from agentforge_core.domain.policy import PolicyDefinition, PolicyValidator +from agentforge_core.guardrails.guardrails_ai import GuardrailsAIEngine +from agentforge_core.llm.mock import MockProvider +from agentforge_core.runtime.orchestrator import AgentOrchestrator + + +def _agent(threshold: int = 4) -> AgentDefinition: + return AgentDefinition( + name="incident_analyzer", + version="v1", + owner="Juan", + purpose="Análisis de incidentes", + state="active", + guardrails=["default"], + llm=LLMConfig(provider="mock"), + system_prompt="Eres un analista. Responde SIEMPRE con JSON.", + output_schema={"type": "object"}, + risk_threshold_for_hitl=threshold, + updated_at=datetime.now(UTC), + ) + + +def _policy(input_validators: list[PolicyValidator] | None = None) -> PolicyDefinition: + return PolicyDefinition( + name="p", + version="v1", + description="t", + input_validators=input_validators or [], + output_validators=[], + ) + + +@pytest.fixture +def orchestrator(tmp_path: Path) -> AgentOrchestrator: + return AgentOrchestrator( + provider=MockProvider(), engine=GuardrailsAIEngine(), data_dir=tmp_path + ) + + +async def test_invoke_camino_feliz(orchestrator: AgentOrchestrator) -> None: + ex = await orchestrator.invoke( + agent_def=_agent(), policy=_policy(), user_input="degradación MOS pool SBC" + ) + assert ex.status == "completed" + assert ex.finished_at is not None + assert ex.needs_human_for is None + assert ex.proposed_actions and ex.proposed_actions[0].action == "scale_out_sbc_pool" + assert ex.final_output is not None and ex.final_output["approved_actions"] + assert [s.step for s in ex.decision_path] == [ + "validate_input", + "llm_reason", + "validate_output", + "propose_actions", + "approve_gate", + "finalize", + ] + + +async def test_invoke_pausa_hitl_y_resume_aprueba(orchestrator: AgentOrchestrator) -> None: + agent = _agent() + paused = await orchestrator.invoke( + agent_def=agent, policy=_policy(), user_input="caída registros sip" + ) + assert paused.status == "awaiting_approval" + assert paused.finished_at is None + assert {a.id for a in paused.needs_human_for or []} == {"act-1"} # act-2 risk=3 < 4 + + resumed = await orchestrator.resume( + agent_def=agent, + policy=_policy(), + trace_id=paused.trace_id, + decision={"approved_action_ids": ["act-1"]}, + ) + assert resumed.status == "completed" + assert resumed.trace_id == paused.trace_id + assert resumed.final_output is not None + assert {a["id"] for a in resumed.final_output["approved_actions"]} == {"act-1"} + + +async def test_resume_rechazo_marca_failed(orchestrator: AgentOrchestrator) -> None: + agent = _agent() + paused = await orchestrator.invoke( + agent_def=agent, policy=_policy(), user_input="caída registros sip" + ) + resumed = await orchestrator.resume( + agent_def=agent, + policy=_policy(), + trace_id=paused.trace_id, + decision={"rejected": True}, + ) + assert resumed.status == "failed" + assert resumed.error == "rejected_by_human" + + +async def test_invoke_bloqueado_por_guardrail(orchestrator: AgentOrchestrator) -> None: + policy = _policy( + [ + PolicyValidator( + type="detect_pii", + config={"entities": ["EMAIL_ADDRESS"], "severity_on_match": "block"}, + ) + ] + ) + ex = await orchestrator.invoke( + agent_def=_agent(), policy=policy, user_input="manda correo a juan@example.com" + ) + assert ex.status == "blocked_by_guardrail" + assert ex.violations and any(v.blocked for v in ex.violations) + assert ex.final_output is None + + +async def test_estado_persiste_entre_instancias(tmp_path: Path) -> None: + agent = _agent() + o1 = AgentOrchestrator( + provider=MockProvider(), engine=GuardrailsAIEngine(), data_dir=tmp_path + ) + paused = await o1.invoke(agent_def=agent, policy=_policy(), user_input="caída registros sip") + assert paused.status == "awaiting_approval" + + # Nueva instancia (simula reinicio del proceso) apuntando al mismo data_dir. + o2 = AgentOrchestrator( + provider=MockProvider(), engine=GuardrailsAIEngine(), data_dir=tmp_path + ) + resumed = await o2.resume( + agent_def=agent, + policy=_policy(), + trace_id=paused.trace_id, + decision={"approved_action_ids": ["act-1"]}, + ) + assert resumed.status == "completed"