foresight / server /agent /graph.py
umangchaudhry's picture
Deploy 1814ed9 from GitHub
2c310c1 verified
Raw
History Blame Contribute Delete
8.88 kB
"""The LangGraph agent and the streaming turn runner.
A deliberately plain ReAct loop: the model calls campus tools until it has enough
to answer, capped so a confused turn can't spin. The interesting parts are around
the edges rather than in the graph shape:
* **No checkpointer.** The durable transcript is the thread JSON in the dataset
repo (see `threads.py`), which the request loads and passes in. A Space sleeps
and restarts, so an in-process checkpointer would be a second, less reliable
source of truth. `compile(checkpointer=...)` is a one-line change if a future
feature needs mid-turn resumption.
* **Citations are derived, not trusted.** Tools record every document they return.
After the turn, the links the model actually wrote are matched back against that
set by URL, so a source chip means "retrieval returned this", not "the model said
so". Anything it linked that retrieval never saw is reported as a web result.
* **Sub-agents are a seam.** The team wants a "plan my next semesters" course
planner later. It becomes another node here; nothing about this file has to
change to accommodate it.
"""
from __future__ import annotations
import logging
import os
import re
from typing import Annotated, TypedDict
from langchain_core.messages import (AIMessage, AIMessageChunk, HumanMessage,
SystemMessage, ToolMessage)
from langgraph.graph import END, START, StateGraph
from langgraph.graph.message import add_messages
from .. import kb
from . import prompts, tools as t
log = logging.getLogger("foresight.agent")
DEFAULT_MODEL = "gpt-5.6-sol"
MAX_TOOL_LOOPS = 6
# Long conversations get expensive and drift. Keep the most recent exchanges;
# the full transcript is always on disk.
MAX_HISTORY_MESSAGES = 24
_MARKDOWN_LINK = re.compile(r"\[[^\]]*\]\((https?://[^\s)]+)\)")
def model_name() -> str:
return os.environ.get("FORESIGHT_CHAT_MODEL", DEFAULT_MODEL)
def enabled() -> bool:
return bool(os.environ.get("OPENAI_API_KEY"))
class State(TypedDict):
messages: Annotated[list, add_messages]
loops: int
def _chat_model():
"""Bound chat model. Imported lazily so the app boots without the SDK or a key."""
from langchain_openai import ChatOpenAI
llm = ChatOpenAI(model=model_name(), use_responses_api=True, streaming=True)
# `web_search` is OpenAI-hosted: it runs server-side, so there's no local
# execution branch for it and results come back already attributed.
return llm.bind_tools([*t.KB_TOOLS, {"type": "web_search"}])
_BY_NAME = {tool.name: tool for tool in t.KB_TOOLS}
def _build():
def agent(state: State):
return {"messages": [_chat_model().invoke(state["messages"])]}
def run_tools(state: State):
last = state["messages"][-1]
out = []
for call in getattr(last, "tool_calls", []) or []:
tool = _BY_NAME.get(call["name"])
if tool is None:
# Hosted tools (web_search) never reach here; anything else is a
# model mistake and should be reported back rather than crash.
out.append(ToolMessage(content=f"Unknown tool: {call['name']}",
tool_call_id=call["id"], name=call["name"]))
continue
try:
result = tool.invoke(call["args"])
except Exception as err:
log.warning("tool %s failed: %s", call["name"], err)
result = {"error": f"{call['name']} failed: {err}"}
out.append(ToolMessage(content=str(result), tool_call_id=call["id"],
name=call["name"]))
return {"messages": out, "loops": state.get("loops", 0) + 1}
def next_step(state: State):
last = state["messages"][-1]
if not getattr(last, "tool_calls", None):
return END
if state.get("loops", 0) >= MAX_TOOL_LOOPS:
log.info("agent: hit the %d-loop cap — answering with what it has", MAX_TOOL_LOOPS)
return END
return "tools"
graph = StateGraph(State)
graph.add_node("agent", agent)
graph.add_node("tools", run_tools)
graph.add_edge(START, "agent")
graph.add_conditional_edges("agent", next_step, {"tools": "tools", END: END})
graph.add_edge("tools", "agent")
return graph.compile()
_graph = None
def get_graph():
global _graph
if _graph is None:
_graph = _build()
return _graph
def _history(messages: list[dict]) -> list:
out = []
for m in messages[-MAX_HISTORY_MESSAGES:]:
text = m.get("text") or ""
if not text:
continue
out.append(HumanMessage(text) if m.get("role") == "user" else AIMessage(text))
return out
def _citations(answer: str, retrieved: dict) -> tuple[list[dict], list[str]]:
"""Split the answer's links into knowledge-base citations and web links.
A source chip means "a tool returned this document during this turn" — so the
only thing matched against is `retrieved`. Deliberately *not* the whole index:
looking a URL up there would hand a chip to a model that guessed a real
vanderbilt.edu address without ever searching for it, which is exactly the
failure the chip is supposed to rule out.
"""
by_url = {d.url: d for d in retrieved.values() if d.url}
cites, web, seen = [], [], set()
for url in _MARKDOWN_LINK.findall(answer):
if url in seen:
continue
seen.add(url)
doc = by_url.get(url)
if doc is not None:
cites.append(doc.cite())
else:
web.append(url)
return cites, web
async def run_turn(question: str, history: list[dict], profile: dict | None,
first_name: str = ""):
"""Run one turn, yielding SSE-shaped events.
Yields dicts: {"type": "tool"|"token"|"sources"|"suggestion"|"error"}.
The caller owns thread ids and persistence.
"""
retrieved, suggestion = t.start_turn()
answer_parts: list[str] = []
announced: set[str] = set()
state = {
"messages": [SystemMessage(prompts.system_prompt(profile, first_name)),
*_history(history), HumanMessage(question)],
"loops": 0,
}
try:
async for kind, payload in get_graph().astream(
state, stream_mode=["messages", "updates"]):
if kind == "messages":
chunk, _meta = payload
# This stream carries every message a node produces, including
# ToolMessages. Only the model's own output is the answer — without
# this check the student watches raw tool JSON scroll past.
if not isinstance(chunk, (AIMessage, AIMessageChunk)):
continue
text = _text_of(chunk)
if text:
answer_parts.append(text)
yield {"type": "token", "text": text}
for call in getattr(chunk, "tool_call_chunks", None) or []:
name = call.get("name")
if name and name not in announced:
announced.add(name)
label = t.TOOL_LABELS.get(name)
if label:
yield {"type": "tool", "name": name, "status": "running",
"label": label}
elif kind == "updates":
for node, update in (payload or {}).items():
if node != "tools":
continue
for msg in update.get("messages", []):
if getattr(msg, "name", None):
yield {"type": "tool", "name": msg.name, "status": "done"}
except Exception as err:
log.exception("agent: turn failed")
yield {"type": "error", "message": str(err)}
return
answer = "".join(answer_parts)
cites, web = _citations(answer, retrieved)
if cites:
yield {"type": "sources", "items": cites}
if web:
yield {"type": "web", "items": web}
if suggestion:
yield {"type": "suggestion", **suggestion}
yield {"type": "final", "text": answer, "sources": cites,
"suggestion": suggestion or None, "tools": sorted(announced)}
def _text_of(chunk) -> str:
"""Text out of a streamed chunk, whose content may be a string or a list of
typed blocks depending on which tools are bound."""
content = getattr(chunk, "content", "")
if isinstance(content, str):
return content
out = []
for block in content or []:
if isinstance(block, str):
out.append(block)
elif isinstance(block, dict) and block.get("type") in ("text", "output_text"):
out.append(block.get("text") or "")
return "".join(out)