ringguard / agent /graph.py
Rohit Patil
RingGuard Phase 3 - initial HF Spaces deploy
753201e
Raw History Blame Contribute Delete
9.11 kB
"""
graph.py
--------
LangGraph orchestration for RingGuard fraud investigation.
State: RingGuardState
cluster — the candidate cluster record
features — extracted feature dict
risk_probability — model score
policy_decision — "auto_hold" | "escalate" | "log_only"
llm_narrative — LLM explanation (None if log_only or if GROQ unavailable)
llm_confidence — LLM confidence statement (None if log_only)
Flow:
score_node → policy_node → [conditional] investigator_node → END
(investigator_node is SKIPPED for log_only decisions)
Usage:
from agent.graph import build_graph, run_cluster
results = run_cluster(cluster, features, risk_probability)
"""
import os
import sys
from typing import Optional, TypedDict
from langgraph.graph import StateGraph, END
BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, BASE_DIR)
from agent.policy import decide as policy_decide
from agent.investigator import investigate
# Audit logger — imported lazily to keep graph.py usable without audit module
_AUDIT_ENABLED = True
try:
from audit.logger import log_case as _log_case
except ImportError:
_AUDIT_ENABLED = False
_log_case = None
# ---------------------------------------------------------------------------
# State definition
# ---------------------------------------------------------------------------
class RingGuardState(TypedDict):
cluster: dict
features: dict
risk_probability: float
policy_decision: Optional[str]
llm_narrative: Optional[str]
llm_confidence: Optional[str]
grounding_ok: Optional[bool]
# ---------------------------------------------------------------------------
# Nodes
# ---------------------------------------------------------------------------
def score_node(state: RingGuardState) -> RingGuardState:
"""
Score node: wraps the already-computed risk_probability from the model.
The probability was computed externally (by scoring/score_demo.py or live
inference) and passed in via state — this node simply forwards it.
"""
# risk_probability is already set in state by the caller
return state
def policy_node(state: RingGuardState) -> RingGuardState:
"""
Policy gate: deterministic decision based on risk_probability + exposure.
Exposure = total_amount from features (Phase 3 will use real Razorpay data).
"""
risk_prob = state["risk_probability"]
exposure = state["features"].get("total_amount", 0.0)
decision = policy_decide(risk_prob, exposure)
return {**state, "policy_decision": decision}
def investigator_node(state: RingGuardState) -> RingGuardState:
"""
LLM investigator: generates a narrative for high-risk clusters.
Only called when policy_decision is 'auto_hold' or 'escalate'.
Catches errors so a single LLM failure does not abort the whole pipeline.
"""
try:
result = investigate(
cluster=state["cluster"],
features=state["features"],
risk_probability=state["risk_probability"],
policy_decision=state["policy_decision"],
)
return {
**state,
"llm_narrative": result["narrative"],
"llm_confidence": result["confidence_statement"],
"grounding_ok": result["grounding_ok"],
}
except Exception as e:
return {
**state,
"llm_narrative": f"[LLM_ERROR] {e}",
"llm_confidence": None,
"grounding_ok": None,
}
# ---------------------------------------------------------------------------
# Conditional routing
# ---------------------------------------------------------------------------
def should_investigate(state: RingGuardState) -> str:
"""
Route to investigator only for auto_hold or escalate decisions.
Skip for log_only — no reason to spend an LLM call on it.
"""
if state.get("policy_decision") in ("auto_hold", "escalate"):
return "investigate"
return "skip"
# ---------------------------------------------------------------------------
# Graph construction
# ---------------------------------------------------------------------------
def build_graph():
"""Build and compile the RingGuard LangGraph state machine."""
builder = StateGraph(RingGuardState)
builder.add_node("score_node", score_node)
builder.add_node("policy_node", policy_node)
builder.add_node("investigator_node", investigator_node)
builder.set_entry_point("score_node")
builder.add_edge("score_node", "policy_node")
builder.add_conditional_edges(
"policy_node",
should_investigate,
{
"investigate": "investigator_node",
"skip": END,
},
)
builder.add_edge("investigator_node", END)
return builder.compile()
# ---------------------------------------------------------------------------
# Convenience runner
# ---------------------------------------------------------------------------
def run_cluster(
cluster: dict,
features: dict,
risk_probability: float,
case_id: str = None,
action_taken: dict = None,
) -> dict:
"""
Run the full pipeline for a single cluster.
Returns final state dict with policy_decision, llm_narrative, etc.
Catches LLM errors and records them in llm_narrative rather than crashing.
Parameters
----------
cluster : candidate cluster record
features : extracted feature dict
risk_probability : model score
case_id : optional case ID for audit logging (e.g. "CASE08")
action_taken : optional Razorpay action dict (populated by Section 3)
"""
graph = build_graph()
initial_state: RingGuardState = {
"cluster": cluster,
"features": features,
"risk_probability": risk_probability,
"policy_decision": None,
"llm_narrative": None,
"llm_confidence": None,
"grounding_ok": None,
}
try:
final_state = graph.invoke(initial_state)
except RuntimeError as e:
# LLM failures (missing key, parse errors) are caught here.
# The policy decision was already set before the LLM call.
# We log the error rather than crashing the whole pipeline.
final_state = {
**initial_state,
"policy_decision": policy_decide(risk_probability, features.get("total_amount", 0.0)),
"llm_narrative": f"[LLM_ERROR] {e}",
"llm_confidence": None,
"grounding_ok": None,
}
result = dict(final_state)
# Phase 3 Section 4.2 — write audit log entry for every run
if _AUDIT_ENABLED and _log_case is not None:
try:
from datetime import datetime, timezone
audit_case = {
"case_id": case_id or cluster.get("cluster_id", "?"),
"cluster_id": cluster.get("cluster_id"),
"customer_ids": cluster.get("customer_ids", []),
"risk_probability": risk_probability,
"policy_decision": result.get("policy_decision"),
"llm_narrative": result.get("llm_narrative"),
"llm_confidence_statement": result.get("llm_confidence"),
"action_taken": action_taken,
"timestamp": datetime.now(timezone.utc).isoformat(),
}
_log_case(audit_case, features)
except Exception:
pass # never crash the pipeline due to audit errors
return result
# ---------------------------------------------------------------------------
# CLI test — run on a single demo cluster
# ---------------------------------------------------------------------------
if __name__ == "__main__":
import json
demo_dir = os.path.join(BASE_DIR, "data", "synthetic")
scored_path = os.path.join(demo_dir, "scored_clusters.json")
orders_path = os.path.join(demo_dir, "orders.json")
graph_path = os.path.join(demo_dir, "graph.gpickle")
with open(scored_path) as f:
scored = json.load(f)
with open(orders_path) as f:
orders = json.load(f)
from scoring.features import extract_features, _load_graph
graph_obj = _load_graph(graph_path) if os.path.exists(graph_path) else None
# Use the first high-risk cluster as a test
test_cluster = scored[0]
features = extract_features(test_cluster, orders, graph=graph_obj)
risk_prob = test_cluster["risk_probability"]
print(f"Testing graph on {test_cluster['cluster_id']} "
f"(risk={risk_prob:.4f})...")
result = run_cluster(test_cluster, features, risk_prob)
print(f" Policy decision : {result['policy_decision']}")
print(f" LLM narrative : {result['llm_narrative']}")
print(f" LLM confidence : {result['llm_confidence']}")
print(f" Grounding OK : {result['grounding_ok']}")