feat(runtime): añade AgentState y nodos LangGraph (validate, llm, propose, approve, finalize)
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,193 @@
|
|||||||
|
"""Factory de nodos LangGraph parametrizados por engine, policy, llm, agent_def."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
from typing import Any
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
import structlog
|
||||||
|
from langgraph.types import interrupt
|
||||||
|
|
||||||
|
from agentforge_core.domain.agent import AgentDefinition
|
||||||
|
from agentforge_core.domain.policy import PolicyDefinition
|
||||||
|
from agentforge_core.guardrails.base import GuardrailEngine
|
||||||
|
from agentforge_core.llm.base import LLMProvider, Message
|
||||||
|
from agentforge_core.runtime.state import AgentState
|
||||||
|
|
||||||
|
log = structlog.get_logger(__name__)
|
||||||
|
|
||||||
|
NodeFn = Callable[[AgentState], Awaitable[dict[str, Any]]]
|
||||||
|
|
||||||
|
|
||||||
|
def _now() -> datetime:
|
||||||
|
return datetime.now(UTC)
|
||||||
|
|
||||||
|
|
||||||
|
def _step(name: str, started: float, **detail: Any) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"step": name,
|
||||||
|
"timestamp": _now().isoformat(),
|
||||||
|
"duration_ms": int((time.perf_counter() - started) * 1000),
|
||||||
|
"detail": detail,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def build_node_validate_input(engine: GuardrailEngine, policy: PolicyDefinition) -> NodeFn:
|
||||||
|
async def validate_input(state: AgentState) -> dict[str, Any]:
|
||||||
|
started = time.perf_counter()
|
||||||
|
trace_id = UUID(state["trace_id"])
|
||||||
|
violations = await engine.validate_input(state["user_input"], policy, trace_id)
|
||||||
|
any_blocked = any(v.blocked for v in violations)
|
||||||
|
new_violations = state.get("violations", []) + [
|
||||||
|
v.model_dump(mode="json") for v in violations
|
||||||
|
]
|
||||||
|
update: dict[str, Any] = {
|
||||||
|
"violations": new_violations,
|
||||||
|
"decision_path": [_step("validate_input", started, n_violations=len(violations))],
|
||||||
|
}
|
||||||
|
if any_blocked:
|
||||||
|
update["status"] = "blocked_by_guardrail"
|
||||||
|
return update
|
||||||
|
|
||||||
|
return validate_input
|
||||||
|
|
||||||
|
|
||||||
|
def build_node_llm_reason(provider: LLMProvider, agent_def: AgentDefinition) -> NodeFn:
|
||||||
|
async def llm_reason(state: AgentState) -> dict[str, Any]:
|
||||||
|
started = time.perf_counter()
|
||||||
|
messages = [
|
||||||
|
Message(role="system", content=agent_def.system_prompt),
|
||||||
|
Message(role="user", content=state["user_input"]),
|
||||||
|
]
|
||||||
|
try:
|
||||||
|
result = await provider.complete(
|
||||||
|
messages=messages,
|
||||||
|
temperature=agent_def.llm.temperature,
|
||||||
|
max_tokens=agent_def.llm.max_tokens,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
log.error("llm_failed", trace_id=state["trace_id"], error=str(exc))
|
||||||
|
return {
|
||||||
|
"status": "failed",
|
||||||
|
"error": "llm_unavailable",
|
||||||
|
"decision_path": [_step("llm_reason", started, ok=False, error=str(exc))],
|
||||||
|
}
|
||||||
|
return {
|
||||||
|
"raw_llm_output": result.content,
|
||||||
|
"messages": [m.model_dump() for m in messages],
|
||||||
|
"decision_path": [
|
||||||
|
_step(
|
||||||
|
"llm_reason",
|
||||||
|
started,
|
||||||
|
model=result.model,
|
||||||
|
tokens_in=result.tokens_in,
|
||||||
|
tokens_out=result.tokens_out,
|
||||||
|
latency_ms=result.latency_ms,
|
||||||
|
)
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
return llm_reason
|
||||||
|
|
||||||
|
|
||||||
|
def build_node_validate_output(engine: GuardrailEngine, policy: PolicyDefinition) -> NodeFn:
|
||||||
|
async def validate_output(state: AgentState) -> dict[str, Any]:
|
||||||
|
started = time.perf_counter()
|
||||||
|
if state.get("status") in {"failed", "blocked_by_guardrail"}:
|
||||||
|
return {}
|
||||||
|
try:
|
||||||
|
parsed = json.loads(state.get("raw_llm_output") or "{}")
|
||||||
|
except json.JSONDecodeError as exc:
|
||||||
|
return {
|
||||||
|
"status": "failed",
|
||||||
|
"error": "output_schema_mismatch",
|
||||||
|
"decision_path": [_step("validate_output", started, ok=False, error=str(exc))],
|
||||||
|
}
|
||||||
|
trace_id = UUID(state["trace_id"])
|
||||||
|
violations = await engine.validate_output(parsed, policy, trace_id)
|
||||||
|
any_blocked = any(v.blocked for v in violations)
|
||||||
|
new_violations = state.get("violations", []) + [
|
||||||
|
v.model_dump(mode="json") for v in violations
|
||||||
|
]
|
||||||
|
update: dict[str, Any] = {
|
||||||
|
"parsed_output": parsed,
|
||||||
|
"violations": new_violations,
|
||||||
|
"decision_path": [_step("validate_output", started, n_violations=len(violations))],
|
||||||
|
}
|
||||||
|
if any_blocked:
|
||||||
|
update["status"] = "blocked_by_guardrail"
|
||||||
|
return update
|
||||||
|
|
||||||
|
return validate_output
|
||||||
|
|
||||||
|
|
||||||
|
def build_node_propose_actions() -> NodeFn:
|
||||||
|
async def propose_actions(state: AgentState) -> dict[str, Any]:
|
||||||
|
started = time.perf_counter()
|
||||||
|
if state.get("status") in {"failed", "blocked_by_guardrail"}:
|
||||||
|
return {}
|
||||||
|
parsed = state.get("parsed_output") or {}
|
||||||
|
actions = parsed.get("proposed_actions", [])
|
||||||
|
return {
|
||||||
|
"proposed_actions": actions,
|
||||||
|
"decision_path": [_step("propose_actions", started, n_actions=len(actions))],
|
||||||
|
}
|
||||||
|
|
||||||
|
return propose_actions
|
||||||
|
|
||||||
|
|
||||||
|
def build_node_approve_gate(agent_def: AgentDefinition) -> NodeFn:
|
||||||
|
async def approve_gate(state: AgentState) -> dict[str, Any]:
|
||||||
|
started = time.perf_counter()
|
||||||
|
if state.get("status") in {"failed", "blocked_by_guardrail"}:
|
||||||
|
return {}
|
||||||
|
actions = state.get("proposed_actions", [])
|
||||||
|
risky = [
|
||||||
|
a
|
||||||
|
for a in actions
|
||||||
|
if int(a.get("risk_score", 1)) >= agent_def.risk_threshold_for_hitl
|
||||||
|
or bool(a.get("requires_approval"))
|
||||||
|
]
|
||||||
|
if not risky:
|
||||||
|
return {"decision_path": [_step("approve_gate", started, hitl=False)]}
|
||||||
|
# Pausa la ejecución hasta que llegue resume(human_decision={...})
|
||||||
|
decision = interrupt({"awaiting_actions": risky})
|
||||||
|
return {
|
||||||
|
"human_decision": decision,
|
||||||
|
"decision_path": [_step("approve_gate", started, hitl=True, resumed=True)],
|
||||||
|
}
|
||||||
|
|
||||||
|
return approve_gate
|
||||||
|
|
||||||
|
|
||||||
|
def build_node_finalize() -> NodeFn:
|
||||||
|
async def finalize(state: AgentState) -> dict[str, Any]:
|
||||||
|
started = time.perf_counter()
|
||||||
|
if state.get("status") in {"failed", "blocked_by_guardrail"}:
|
||||||
|
return {"decision_path": [_step("finalize", started, skipped=True)]}
|
||||||
|
parsed = state.get("parsed_output") or {}
|
||||||
|
decision = state.get("human_decision") or {}
|
||||||
|
approved_ids = set(decision.get("approved_action_ids", [])) if decision else None
|
||||||
|
actions = state.get("proposed_actions", [])
|
||||||
|
if approved_ids is not None:
|
||||||
|
final_actions = [a for a in actions if a.get("id") in approved_ids]
|
||||||
|
if decision.get("rejected"):
|
||||||
|
return {
|
||||||
|
"status": "failed",
|
||||||
|
"error": "rejected_by_human",
|
||||||
|
"decision_path": [_step("finalize", started, rejected=True)],
|
||||||
|
}
|
||||||
|
else:
|
||||||
|
final_actions = actions
|
||||||
|
final_output = {**parsed, "approved_actions": final_actions}
|
||||||
|
return {
|
||||||
|
"final_output": final_output,
|
||||||
|
"status": "completed",
|
||||||
|
"decision_path": [_step("finalize", started, n_approved=len(final_actions))],
|
||||||
|
}
|
||||||
|
|
||||||
|
return finalize
|
||||||
@@ -0,0 +1,25 @@
|
|||||||
|
"""Estado tipado del grafo LangGraph."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import operator
|
||||||
|
from typing import Annotated, Any, TypedDict
|
||||||
|
|
||||||
|
|
||||||
|
class AgentState(TypedDict, total=False):
|
||||||
|
"""Estado mutable que fluye por el grafo; `decision_path` se acumula."""
|
||||||
|
|
||||||
|
trace_id: str
|
||||||
|
agent_name: str
|
||||||
|
agent_version: str
|
||||||
|
user_input: str
|
||||||
|
messages: list[dict[str, Any]]
|
||||||
|
raw_llm_output: str | None
|
||||||
|
parsed_output: dict[str, Any] | None
|
||||||
|
proposed_actions: list[dict[str, Any]]
|
||||||
|
violations: list[dict[str, Any]]
|
||||||
|
decision_path: Annotated[list[dict[str, Any]], operator.add]
|
||||||
|
status: str
|
||||||
|
error: str | None
|
||||||
|
human_decision: dict[str, Any] | None
|
||||||
|
final_output: dict[str, Any] | None
|
||||||
@@ -0,0 +1,66 @@
|
|||||||
|
"""Tests de los nodos del grafo (lógica pura, sin compilación de grafo)."""
|
||||||
|
|
||||||
|
from uuid import uuid4
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from agentforge_core.domain.policy import PolicyDefinition, PolicyValidator
|
||||||
|
from agentforge_core.guardrails.guardrails_ai import GuardrailsAIEngine
|
||||||
|
from agentforge_core.runtime.nodes import build_node_validate_input
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def engine() -> GuardrailsAIEngine:
|
||||||
|
return GuardrailsAIEngine()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def policy_block_email() -> PolicyDefinition:
|
||||||
|
return PolicyDefinition(
|
||||||
|
name="t",
|
||||||
|
version="v1",
|
||||||
|
description="t",
|
||||||
|
input_validators=[
|
||||||
|
PolicyValidator(
|
||||||
|
type="detect_pii",
|
||||||
|
config={"entities": ["EMAIL_ADDRESS"], "severity_on_match": "block"},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
output_validators=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _base_state(user_input: str) -> dict:
|
||||||
|
return {
|
||||||
|
"trace_id": str(uuid4()),
|
||||||
|
"agent_name": "x",
|
||||||
|
"agent_version": "v1",
|
||||||
|
"user_input": user_input,
|
||||||
|
"messages": [],
|
||||||
|
"raw_llm_output": None,
|
||||||
|
"parsed_output": None,
|
||||||
|
"proposed_actions": [],
|
||||||
|
"violations": [],
|
||||||
|
"decision_path": [],
|
||||||
|
"status": "running",
|
||||||
|
"error": None,
|
||||||
|
"human_decision": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_validate_input_marca_blocked_si_pii(
|
||||||
|
engine: GuardrailsAIEngine, policy_block_email: PolicyDefinition
|
||||||
|
) -> None:
|
||||||
|
node = build_node_validate_input(engine, policy_block_email)
|
||||||
|
out = await node(_base_state("manda correo a juan@example.com"))
|
||||||
|
assert out["status"] == "blocked_by_guardrail"
|
||||||
|
assert any(v["blocked"] for v in out["violations"])
|
||||||
|
|
||||||
|
|
||||||
|
async def test_validate_input_pasa_sin_pii(
|
||||||
|
engine: GuardrailsAIEngine, policy_block_email: PolicyDefinition
|
||||||
|
) -> None:
|
||||||
|
node = build_node_validate_input(engine, policy_block_email)
|
||||||
|
out = await node(_base_state("incidente sin pii"))
|
||||||
|
# El nodo solo escribe `status` cuando bloquea; si no, el estado sigue "running".
|
||||||
|
assert out.get("status", "running") == "running"
|
||||||
Reference in New Issue
Block a user