Skip to content

Commit 834ee0b

Browse files
added graph generator guardrails
1 parent f7d0368 commit 834ee0b

3 files changed

Lines changed: 341 additions & 42 deletions

File tree

app/agent/graph.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,10 +44,18 @@ def call_model(state: GraphState):
4444
"CRITICAL RULES:\n"
4545
"- If the user asks for a chart or graph, you MUST call 'generate_graph'."
4646
" Do NOT just describe data in text.\n"
47+
"- Never print JSON like {\"name\": \"generate_graph\", ...} in normal text."
48+
" Use real tool-calling only.\n"
4749
"- First use 'search_documents' to gather the numbers if needed,"
4850
" then call 'generate_graph' with the data.\n"
51+
"- Only chart numbers that are explicitly present in retrieved document evidence."
52+
" Never invent placeholder values.\n"
53+
"- For requests like monthly revenue for a specific year, include all 12 months"
54+
" if the evidence supports it; otherwise clearly state the data is incomplete.\n"
4955
"- Never call 'search_documents' repeatedly with the same query in a loop."
5056
" If you already called it and have results, produce a final answer.\n"
57+
"- If 'generate_graph' returns an error, do NOT call it again with the same payload."
58+
" Return a concise failure message asking for a clearer chart request.\n"
5159
"- After calling 'generate_graph', keep your final text response brief (e.g. 'Here is your chart.'). "
5260
"Do NOT re-list all the data points in your text answer — the chart already shows them."
5361
)

app/agent/tools/graph_generator.py

Lines changed: 62 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
"""LangGraph tool: generate Plotly chart payloads from structured JSON."""
22

3+
import ast
34
import json
5+
import re
46
from typing import Any
57

68
from langchain_core.tools import tool
@@ -10,8 +12,59 @@
1012
logger = get_logger(__name__)
1113

1214

15+
def _parse_graph_input(data_json: Any) -> dict[str, Any]:
16+
"""Parse graph payload from strict JSON or common dict-like string variants."""
17+
if isinstance(data_json, dict):
18+
return data_json
19+
20+
if not isinstance(data_json, str):
21+
raise ValueError("Graph input must be a JSON string or dictionary.")
22+
23+
content = data_json.strip()
24+
if not content:
25+
raise ValueError("Graph input is empty.")
26+
27+
# Some model outputs wrap JSON in markdown code fences.
28+
if content.startswith("```"):
29+
content = re.sub(r"^```(?:json)?\\s*", "", content)
30+
content = re.sub(r"\\s*```$", "", content).strip()
31+
32+
try:
33+
parsed = json.loads(content)
34+
if isinstance(parsed, dict):
35+
return parsed
36+
except json.JSONDecodeError:
37+
pass
38+
39+
# Fallback for Python dict-style strings (single quotes / True / False / None).
40+
try:
41+
literal = ast.literal_eval(content)
42+
if isinstance(literal, dict):
43+
return literal
44+
except (SyntaxError, ValueError):
45+
pass
46+
47+
# Last attempt: extract the first object-like fragment and parse it.
48+
match = re.search(r"\{.*\}", content, re.DOTALL)
49+
if match:
50+
fragment = match.group(0)
51+
try:
52+
parsed = json.loads(fragment)
53+
if isinstance(parsed, dict):
54+
return parsed
55+
except json.JSONDecodeError:
56+
try:
57+
literal = ast.literal_eval(fragment)
58+
if isinstance(literal, dict):
59+
return literal
60+
except (SyntaxError, ValueError):
61+
pass
62+
63+
raise ValueError("Invalid graph payload format.")
64+
65+
1366
@tool
14-
def generate_graph(data_json: str) -> str:
67+
def generate_graph(data_json: Any) -> str:
1568
"""
1669
Generate an interactive chart or graph in Plotly JSON format.
1770
Use this tool ONLY when the user explicitly asks for a chart, graph, plot, or visual breakdown.
@@ -30,13 +83,17 @@ def generate_graph(data_json: str) -> str:
3083
'{"title": "Q1 Revenue", "chart_type": "bar",
3184
"labels": ["Jan", "Feb", "Mar"], "values": [100, 150, 200]}'
3285
"""
33-
logger.info("Graph generation invoked", extra={"payload_length": len(data_json)})
86+
payload_length = len(data_json) if isinstance(data_json, str) else None
87+
logger.info("Graph generation invoked", extra={"payload_length": payload_length, "input_type": type(data_json).__name__})
3488

3589
try:
36-
data = json.loads(data_json)
37-
except json.JSONDecodeError as exc:
90+
data = _parse_graph_input(data_json)
91+
except ValueError as exc:
3892
logger.warning("Graph generation failed — invalid JSON input", extra={"reason": str(exc)})
39-
return json.dumps({"error": "Invalid JSON provided. Please provide a valid JSON string."})
93+
return json.dumps({
94+
"error": "Invalid graph payload. Provide a JSON object with title, chart_type, labels, and values/series.",
95+
"do_not_retry": True,
96+
})
4097

4198
title = data.get("title", "Chart")
4299
chart_type = data.get("chart_type", "bar")

0 commit comments

Comments
 (0)