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:
Juan
2026-05-10 10:56:25 +02:00
co-authored by Claude Opus 4.7
parent 006d781df0
commit 683cc69edc
3 changed files with 284 additions and 0 deletions
+193
View File
@@ -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
+25
View File
@@ -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