File size: 8,883 Bytes
2c310c1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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)