File size: 17,874 Bytes
b8f211a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
af2544e
 
 
 
b8f211a
 
 
 
 
 
 
 
 
 
 
 
af2544e
b8f211a
af2544e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a26ee2c
 
 
 
 
 
 
 
 
 
8787588
 
 
 
 
 
 
 
a26ee2c
8787588
 
a26ee2c
 
8787588
 
a26ee2c
 
 
 
 
 
 
 
 
 
af2544e
 
 
 
 
 
 
b8f211a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
af2544e
 
 
 
 
 
 
 
 
 
 
 
 
 
b8f211a
 
 
 
 
 
af2544e
 
 
 
 
 
b8f211a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
af2544e
 
b8f211a
 
 
 
 
 
 
 
af2544e
b8f211a
 
af2544e
b8f211a
 
 
 
af2544e
b8f211a
 
 
 
 
af2544e
b8f211a
 
 
 
 
 
 
 
 
 
 
af2544e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b8f211a
af2544e
 
b8f211a
af2544e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
82732c2
af2544e
 
b8f211a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
af2544e
 
 
b8f211a
af2544e
 
b8f211a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
964ad3d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
82732c2
964ad3d
 
 
 
 
 
b8f211a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
964ad3d
b8f211a
 
af2544e
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
"""
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