nagent / agents.py
techresearchspace's picture
Update agents.py
82732c2 verified
Raw History Blame Contribute Delete
17.9 kB
"""
Multi-agent orchestration with LangGraph.
Explicit state machine: a manager node decides which specialist to route to
next, each specialist does its work and returns to the manager, and the
manager decides when to finish. This gives you deterministic control flow
instead of framework-decided delegation.
"""
import os
import json
import subprocess
import sys
import tempfile
from typing import Annotated, Literal, TypedDict
import spaces
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from langchain_core.messages import AnyMessage, AIMessage, HumanMessage, SystemMessage
from langchain_core.tools import tool
from langchain_tavily import TavilySearch
from langgraph.graph import StateGraph, END
from langgraph.graph.message import add_messages
# Hard ceiling on manager<->specialist round-trips, independent of the graph's
# recursion_limit. If the manager keeps flip-flopping between specialists
# without ever saying "finish", this forces finalize instead of erroring out.
MAX_TURNS = 6
# ---------------------------------------------------------------------------
# Local model, run on a ZeroGPU-allocated GPU
# ---------------------------------------------------------------------------
# Pick a model that fits comfortably in a ZeroGPU allocation. 3B is a safe
# starting point; move up to 7B/8B if quality is the bottleneck, but watch
# the `duration` below β€” bigger models need more time per call.
MODEL_ID = "Qwen/Qwen2.5-3B-Instruct"
# Loaded once at import time, kept on CPU. ZeroGPU only attaches a GPU for
# the duration of a function decorated with @spaces.GPU β€” the model is
# moved onto that GPU inside `generate()` below and this is the standard
# pattern for ZeroGPU Spaces (see huggingface.co/docs/hub/spaces-zerogpu).
_tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
_model = AutoModelForCausalLM.from_pretrained(MODEL_ID, torch_dtype=torch.bfloat16)
@spaces.GPU(duration=60)
def generate(messages: list[dict], max_new_tokens: int = 512) -> str:
"""Run one generation on a ZeroGPU-allocated GPU. `messages` is a plain
list of {"role": ..., "content": ...} dicts (chat-template format).
Any exception is caught HERE, inside the GPU worker, rather than left to
propagate across the @spaces.GPU process boundary β€” exceptions crossing
that boundary can come back mangled (e.g. a real AttributeError showing
up as an opaque KeyError('AttributeError') in the caller). Catching and
returning a plain string keeps the real error message intact and visible.
"""
try:
_model.to("cuda")
# return_dict=True guarantees a BatchEncoding with named fields
# (input_ids, attention_mask) rather than leaving it to whatever
# the installed transformers version defaults to β€” some versions
# return a plain tensor here, others a BatchEncoding, and treating
# the latter as a tensor (e.g. calling .shape on it directly) fails
# with an unhelpful bare AttributeError.
encoded = _tokenizer.apply_chat_template(
messages, add_generation_prompt=True, return_tensors="pt", return_dict=True
).to("cuda")
input_ids = encoded["input_ids"]
attention_mask = encoded.get("attention_mask")
with torch.no_grad():
output_ids = _model.generate(
input_ids=input_ids,
attention_mask=attention_mask,
max_new_tokens=max_new_tokens,
do_sample=False,
)
new_tokens = output_ids[0][input_ids.shape[-1]:]
return _tokenizer.decode(new_tokens, skip_special_tokens=True).strip()
except Exception as e:
import traceback
tb = traceback.format_exc()
print(f"[generate] {type(e).__name__}: {e}\n{tb}", flush=True) # full detail in container logs
return f'{{"error": "{type(e).__name__}: {e}"}}'
def _lc_to_dicts(messages: list[AnyMessage]) -> list[dict]:
"""Convert langchain message objects into the plain role/content dicts
apply_chat_template expects."""
role_map = {"human": "user", "ai": "assistant", "system": "system"}
return [{"role": role_map.get(m.type, "user"), "content": m.content} for m in messages]
# ---------------------------------------------------------------------------
# Shared graph state
# ---------------------------------------------------------------------------
class AgentState(TypedDict):
messages: Annotated[list[AnyMessage], add_messages]
next: str # which node the manager wants to run next
task: str # original user task, kept for reference
results: dict # scratch space specialists write into
turn_count: int # manager round-trips so far, enforces MAX_TURNS
# ---------------------------------------------------------------------------
# Tools per specialist
# ---------------------------------------------------------------------------
# Requires TAVILY_API_KEY as a Space secret. Tavily is built for LLM agents
# (returns clean, citeable snippets rather than raw HTML), which is why it's
# preferred here over scraping a generic search engine yourself.
#
# Lazily constructed on first use rather than at import time: TavilySearch
# validates its API key eagerly in __init__, so building it at module load
# means a missing/bad key crashes the entire app before Gradio even starts.
# Deferring it means only the search tool fails, with a clear message,
# while the rest of the app (code_agent, data_agent) still works.
_tavily_instance = None
def _get_tavily():
global _tavily_instance
if _tavily_instance is None:
_tavily_instance = TavilySearch(max_results=5)
return _tavily_instance
@tool
def web_search(query: str) -> str:
"""Search the web for a query and return a short summary of results."""
try:
tavily = _get_tavily()
except Exception as e:
return f"Search unavailable: TAVILY_API_KEY missing or invalid ({e})"
try:
results = tavily.invoke({"query": query})
except Exception as e:
return f"Search failed: {e}"
items = results.get("results", []) if isinstance(results, dict) else results
if not items:
return "No results found."
lines = []
for r in items[:5]:
title = r.get("title", "")
url = r.get("url", "")
content = (r.get("content", "") or "")[:300]
lines.append(f"- {title} ({url}): {content}")
return "\n".join(lines)
# Sandboxing note: this uses a locked-down subprocess as a reasonable
# starting point (no shell, no filesystem/network access from the executed
# code's perspective within the limits below, hard timeout, restricted
# builtins). It is NOT a substitute for real isolation. For production,
# run this in a dedicated sandbox with its own OS-level boundary β€” e.g.
# E2B Code Interpreter, Modal Sandboxes, or a locked-down gVisor/Firecracker
# container β€” rather than a subprocess on the Space's own host.
_ALLOWED_BUILTIN_NAMES = [
"print", "range", "len", "str", "int", "float", "bool", "list",
"dict", "set", "tuple", "enumerate", "zip", "map", "filter",
"sum", "min", "max", "sorted", "abs", "round",
]
# This template is written to a temp file and run as a separate, isolated
# (`-I`) interpreter process. It builds its OWN restricted builtins dict at
# runtime inside that fresh process β€” it does not reuse anything from this
# module's environment.
_SANDBOX_RUNNER = """
import sys
import builtins as _builtins_module
code = sys.stdin.read()
allowed = %r
safe_builtins = {name: getattr(_builtins_module, name) for name in allowed if hasattr(_builtins_module, name)}
env = {"__builtins__": safe_builtins}
try:
exec(code, env)
except Exception as e:
print(f"Error: {e}", file=sys.stderr)
sys.exit(1)
""" % (_ALLOWED_BUILTIN_NAMES,)
@tool
def run_python(code: str) -> str:
"""Execute a snippet of Python in a restricted subprocess and return stdout or the error."""
runner_src = _SANDBOX_RUNNER
try:
with tempfile.NamedTemporaryFile("w", suffix=".py", delete=False) as f:
f.write(runner_src)
runner_path = f.name
proc = subprocess.run(
[sys.executable, "-I", runner_path], # -I: isolated mode, ignores env/site
input=code,
capture_output=True,
text=True,
timeout=5,
)
if proc.returncode != 0:
return f"Error: {proc.stderr.strip() or 'execution failed'}"
return proc.stdout or "(no output)"
except subprocess.TimeoutExpired:
return "Error: execution timed out after 5s"
except Exception as e:
return f"Error: {e}"
finally:
try:
os.remove(runner_path)
except OSError:
pass
@tool
def summarize_numbers(numbers: list[float]) -> str:
"""Compute summary statistics (count, sum, mean, min, max) for numbers."""
if not numbers:
return "No numbers provided."
n = len(numbers)
total = sum(numbers)
return f"count={n}, sum={total:.2f}, mean={total / n:.2f}, min={min(numbers)}, max={max(numbers)}"
SEARCH_TOOLS = [web_search]
CODE_TOOLS = [run_python]
DATA_TOOLS = [summarize_numbers]
# ---------------------------------------------------------------------------
# Manager node β€” decides which specialist runs next, or "finish"
# ---------------------------------------------------------------------------
MANAGER_SYSTEM_PROMPT = """You are a manager coordinating three specialist agents:
- search_agent: looks things up on the web
- code_agent: writes and runs Python
- data_agent: computes statistics on numeric data
Given the conversation so far, decide the single next step. Respond with ONLY
a JSON object of the form: {"next": "search_agent" | "code_agent" | "data_agent" | "finish", "reason": "..."}
Choose "finish" once you have enough information to answer the user directly.
"""
def manager_node(state: AgentState):
turn_count = state.get("turn_count", 0) + 1
# Hard stop: don't even call the model once we've hit the cap β€” force
# finish so a manager stuck in a routing loop can't burn turns indefinitely.
if turn_count > MAX_TURNS:
return {
"next": "finish",
"turn_count": turn_count,
"messages": [SystemMessage(content=f"Max turns ({MAX_TURNS}) reached, finalizing with available results.")],
}
messages = [SystemMessage(content=MANAGER_SYSTEM_PROMPT)] + state["messages"]
raw_response = generate(_lc_to_dicts(messages), max_new_tokens=200)
try:
raw = raw_response.strip().strip("`")
if raw.startswith("json"):
raw = raw[4:]
decision = json.loads(raw)
next_step = decision.get("next", "finish")
except json.JSONDecodeError:
next_step = "finish"
if next_step not in {"search_agent", "code_agent", "data_agent", "finish"}:
next_step = "finish"
return {"next": next_step, "turn_count": turn_count, "messages": [AIMessage(content=raw_response)]}
def route_from_manager(state: AgentState) -> Literal["search_agent", "code_agent", "data_agent", "finalize"]:
if state["next"] == "finish":
return "finalize"
return state["next"]
# ---------------------------------------------------------------------------
# Specialist nodes β€” each does its work, appends a message, returns to manager
# ---------------------------------------------------------------------------
def _tool_spec_text(tools) -> str:
lines = []
for t in tools:
lines.append(f'- {t.name}: {t.description} Args schema: {t.args}')
return "\n".join(lines)
def make_specialist_node(name: str, tools, role_description: str):
tool_spec = _tool_spec_text(tools)
system_prompt = (
f"{role_description}\n\n"
f"Available tools:\n{tool_spec}\n\n"
'Respond with ONLY a JSON object. Either call a tool: '
'{"tool": "<tool_name>", "args": {...}} β€” or, if no tool is needed, '
'give the answer directly: {"final": "<your answer>"}'
)
def node(state: AgentState):
messages = [{"role": "system", "content": system_prompt}, {"role": "user", "content": state["task"]}]
raw_response = generate(messages, max_new_tokens=300)
summary = None
try:
raw = raw_response.strip().strip("`")
if raw.startswith("json"):
raw = raw[4:]
decision = json.loads(raw)
except json.JSONDecodeError:
decision = None
if decision and "tool" in decision:
tool_fn = next((t for t in tools if t.name == decision["tool"]), None)
if tool_fn is not None:
try:
tool_output = tool_fn.invoke(decision.get("args", {}))
except Exception as e:
tool_output = f"Tool error: {e}"
# One follow-up call to turn the raw tool output into a
# concise summary, rather than dumping it unprocessed.
followup = messages + [
{"role": "assistant", "content": raw_response},
{"role": "user", "content": f"Tool result: {tool_output}\n\nSummarize this for the manager."},
]
summary = generate(followup, max_new_tokens=250)
else:
summary = f"Requested unknown tool '{decision['tool']}'."
elif decision and "final" in decision:
summary = str(decision["final"]) # model may return a bare number/bool instead of a string
else:
summary = raw_response # model didn't follow the JSON protocol β€” use raw text as a fallback
state["results"][name] = summary
return {
"messages": [HumanMessage(content=f"[{name}] {summary}")],
"results": state["results"],
}
return node
search_node = make_specialist_node(
"search_agent", SEARCH_TOOLS, "You answer questions using the web_search tool. Be concise."
)
code_node = make_specialist_node(
"code_agent", CODE_TOOLS, "You solve tasks by writing and running Python with the run_python tool."
)
data_node = make_specialist_node(
"data_agent", DATA_TOOLS, "You analyze numeric data using the summarize_numbers tool."
)
# ---------------------------------------------------------------------------
# Finalize node β€” synthesizes a final answer from gathered specialist results
# ---------------------------------------------------------------------------
def finalize_node(state: AgentState):
context = "\n".join(f"{k}: {v}" for k, v in state["results"].items())
messages = [
{"role": "system", "content": "Write a final answer for the user based on the specialist findings below."},
{"role": "user", "content": f"Task: {state['task']}\n\nFindings:\n{context}"},
]
response = generate(messages, max_new_tokens=500)
return {"messages": [AIMessage(content=response)]}
# ---------------------------------------------------------------------------
# Build the graph
# ---------------------------------------------------------------------------
def build_graph():
graph = StateGraph(AgentState)
graph.add_node("manager", manager_node)
graph.add_node("search_agent", search_node)
graph.add_node("code_agent", code_node)
graph.add_node("data_agent", data_node)
graph.add_node("finalize", finalize_node)
graph.set_entry_point("manager")
graph.add_conditional_edges(
"manager",
route_from_manager,
{
"search_agent": "search_agent",
"code_agent": "code_agent",
"data_agent": "data_agent",
"finalize": "finalize",
},
)
# each specialist reports back to the manager for the next decision
graph.add_edge("search_agent", "manager")
graph.add_edge("code_agent", "manager")
graph.add_edge("data_agent", "manager")
graph.add_edge("finalize", END)
return graph.compile()
AGENT_LABELS = {
"search_agent": "πŸ” Search agent",
"code_agent": "πŸ’» Code agent",
"data_agent": "πŸ“Š Data agent",
}
def _format_attribution(results: dict) -> str:
"""Build a deterministic 'which agent did what' footer from state.results.
Built in code rather than asked of the model, so attribution can't drift
from what actually ran β€” dict insertion order matches invocation order
since specialists write into results as they're called."""
if not results:
return ""
lines = ["", "---", "**Agents used:**"]
for name, summary in results.items():
summary = str(summary) # defensive: a tool or model could hand back a non-string value
label = AGENT_LABELS.get(name, name)
preview = summary if len(summary) <= 160 else summary[:157] + "..."
lines.append(f"- **{label}** β€” {preview}")
return "\n".join(lines)
def run_task(app, task: str, session_state: dict) -> str:
"""Run one task through the graph. session_state is the gr.State dict
used to persist conversation history across turns in the UI."""
initial_state: AgentState = {
"messages": [HumanMessage(content=task)],
"next": "",
"task": task,
"results": {},
"turn_count": 0,
}
# recursion_limit is a blunt backstop (counts every node hop, including
# finalize); MAX_TURNS in manager_node is the real, tuned cap on manager
# round-trips and is what actually prevents runaway loops.
final_state = app.invoke(initial_state, config={"recursion_limit": 2 * MAX_TURNS + 5})
answer = final_state["messages"][-1].content
answer += _format_attribution(final_state.get("results", {}))
session_state.setdefault("history", []).append({"task": task, "answer": answer})
return answer