Download agent/graph.py from OGrohit/ringguard: direct link, hf CLI and curl.
- Browser
- Download file 9.11 kB
-
https://huggingface.co/spaces/OGrohit/ringguard/resolve/main/agent/graph.py
- Command line
-
hf download hf://spaces/OGrohit/ringguard/agent/graph.py
-
curl -L -o graph.py https://huggingface.co/spaces/OGrohit/ringguard/resolve/main/agent/graph.py
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']}") | |