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