File size: 10,795 Bytes
b2931f4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""The agent's LangGraph nodes.

Flow: rewrite_query β†’ route β†’ (retrieve?) β†’ agent(tool-loop+synthesis) β†’ END.

Design note: the handoff listed tool-loop and synthesize as separate steps, but
with native function-calling they're one node by construction β€” the loop runs
until the model stops emitting function_calls and produces its final text, and
that terminal text *is* the synthesis. We still emit a distinct 'synthesize'
trace event so the frontend (Decision 18) can render it as its own step.
"""

from __future__ import annotations

from functools import lru_cache

from finrag.agent.state import AgentState
from finrag.ingestion.facts import corpus_companies, corpus_years
from finrag.llm import generate_text, run_tool_loop_stream, synthesize
from finrag.llm.base import ToolCall, format_chunks_for_prompt
from finrag.retrieval.rerank import rerank_search

MAX_TOOL_ITERS = 5  # hard cap so a confused model can't loop forever


@lru_cache(maxsize=1)
def _corpus_grounding() -> str:
    """Tell the model the exact known universe so vague references ('these three
    companies', 'all of them') resolve to real corpus members instead of the
    model guessing (it would otherwise pull in Microsoft/Google). Data-driven β€”
    reads the loaded DuckDB, so it can never drift from what's queryable."""
    companies = corpus_companies()
    if not companies:  # corpus not loaded β€” emit nothing rather than a wrong claim
        return ""
    lo, hi = corpus_years()
    span = f"fiscal years {lo}–{hi}" if lo and hi else "the available fiscal years"
    listing = "; ".join(f"{name} ({ticker})" for ticker, name in companies)
    return (
        f"\n\nKNOWN CORPUS β€” the dataset contains EXACTLY these {len(companies)} "
        f"companies, {span}: {listing}.\n"
        "When the question refers to the companies without naming them ('these "
        "companies', 'the three companies', 'all of them', 'each company'), it means "
        "exactly this set β€” resolve the reference to these names. Never introduce a "
        "company outside this set; if asked about one that isn't listed, say it is not "
        "in the corpus rather than answering from general knowledge."
    )


def _stream_writer():
    """LangGraph custom-stream writer when the graph is driven by
    `graph.stream(..., stream_mode=[..., "custom"])` (the /agent/stream SSE
    path); a no-op otherwise (plain invoke / direct call). The same agent node
    therefore serves both /agent and /agent/stream without branching."""
    try:
        from langgraph.config import get_stream_writer

        return get_stream_writer()
    except Exception:
        return lambda _data: None

# ── Prompts ────────────────────────────────────────────────────────────────
# rewrite + route merged into ONE call to save a request against the free-tier
# 5-req/min cap (the agent is call-heavy). The route hint is also tightened:
# segment-level figures (services/product revenue) live in narrative, not in
# our top-level XBRL facts, so they must route to vector β€” this fixes the
# earlier mis-route that answered "services revenue" with total revenue.
_PLAN_SYSTEM = """You prepare a question about SEC 10-K filings for retrieval. Do two things:

1. Rewrite it as a concise, self-contained search query: resolve vague references, and make the company and fiscal year explicit if implied.
2. Classify what it needs:
   - vector : qualitative/narrative content, OR segment-level figures like services/product/regional revenue (these live in the filing text, not the figures database)
   - sql    : precise TOP-LEVEL financials (total revenue, net income, total assets, margins, multi-year or cross-company comparisons)
   - both   : needs narrative AND exact top-level figures

Output EXACTLY two lines, nothing else:
QUERY: <rewritten query>
ROUTE: <vector|sql|both>"""

_AGENT_SYSTEM = """You are a financial analyst assistant answering questions about SEC 10-K filings.

You have tools:
- sql_query: get EXACT figures for TOP-LEVEL metrics only (total revenue, net income, total assets, margins). Prefer it for those over reading numbers from text. It does NOT have segment/product/regional figures (e.g. services revenue, iPhone revenue) β€” for those, read the value from the context chunks and cite [N]. If sql_query returns an error, fall back to the context.
- calculator: do arithmetic (growth rates, margins, ratios). Extract numbers, then compute β€” never do multi-digit math in your head.
- lookup_citation: re-fetch a chunk's full text by chunk_id if you need to quote it exactly. Pass the exact id shown as (id=...) in the chunk's header β€” never the [N] anchor.

Rules:
1. Ground every claim in the provided context chunks or tool results. If neither contains the answer, say so β€” do not use prior knowledge.
2. A figure must actually match what was asked. If a tool returns a number for a different metric than the question, do not report it β€” use the context instead.
3. Cite narrative facts from the context with [N], where N is the chunk index shown. Cite even when paraphrasing.
4. Quote exact figures; never round unless asked.
5. Be careful with fiscal vs calendar year (Apple's fiscal year ends in late September).
6. Be concise β€” match the question's scope.
7. Do not embellish. State only what the context or tool results actually support. Do not add provenance you cannot see (e.g. "as disclosed in the 10-K" when the figure came from the figures database), characterizations ("a record-setting profit", "strong performance"), or outside facts not present in the context/tool results. A correct figure with ungrounded commentary is still a faithfulness failure.
"""


def _parse_plan(raw: str, fallback_query: str) -> tuple[str, str]:
    """Parse the two-line plan output into (rewritten_query, route)."""
    query, route = fallback_query, "both"
    for line in raw.splitlines():
        s = line.strip()
        low = s.lower()
        if low.startswith("query:"):
            query = s.split(":", 1)[1].strip() or fallback_query
        elif low.startswith("route:"):
            r = s.split(":", 1)[1].strip().lower()
            if "both" in r:
                route = "both"
            elif "sql" in r:
                route = "sql"
            elif "vector" in r:
                route = "vector"
    return query, route


def plan(state: AgentState) -> AgentState:
    """One call that both rewrites the query and routes it. Emits two trace
    events so the frontend still shows rewrite and route as distinct steps."""
    original = state["question"]
    # Ground the rewrite in the known corpus so "these 3 companies" expands to the
    # real names here, before retrieval and the agent ever see the query.
    raw = generate_text(_PLAN_SYSTEM + _corpus_grounding(), original, max_output_tokens=128)
    rewritten, decision = _parse_plan(raw, original)
    return {
        "rewritten_query": rewritten,
        "route": decision,
        "trace": [
            {
                "node": "plan",
                "type": "rewrite",
                "data": {"original": original, "rewritten": rewritten},
            },
            {"node": "plan", "type": "route", "data": {"route": decision}},
        ],
    }


def retrieve(state: AgentState) -> AgentState:
    """Vector retrieval via the Day-2 funnel. Reached only when the route
    includes vector (conditional edge in graph.py)."""
    chunks = rerank_search(question=state["rewritten_query"], top_k=8)
    return {
        "chunks": chunks,
        "trace": [
            {
                "node": "retrieve",
                "type": "retrieve",
                "data": {
                    "n_chunks": len(chunks),
                    "top": [
                        {"chunk_id": c.chunk_id, "ticker": c.ticker, "fy": c.fiscal_year}
                        for c in chunks[:3]
                    ],
                },
            }
        ],
    }


def agent(state: AgentState) -> AgentState:
    """Provider-neutral tool-calling loop (Claude tool_use or Gemini function
    calling, per llm_provider). Runs tools until a final text answer, surfacing
    each tool call in the trace (the SQL/args are the demo payload)."""
    chunks = state.get("chunks", [])
    context = (
        format_chunks_for_prompt(chunks)
        if chunks
        else "(no vector context retrieved β€” rely on tools)"
    )
    user_text = f"Question: {state['rewritten_query']}\n\nContext chunks:\n\n{context}"

    # Push live events to the SSE stream (no-op under plain /agent). Tokens are
    # the final answer forming; tool_call fires the instant a tool runs.
    writer = _stream_writer()

    def on_text(delta: str) -> None:
        writer({"type": "token", "text": delta})

    def on_tool_call(tc: ToolCall) -> None:
        writer(
            {
                "type": "tool_call",
                "node": "agent",
                "data": {"tool": tc.tool, "args": tc.args, "result": tc.result},
            }
        )

    result = run_tool_loop_stream(
        _AGENT_SYSTEM + _corpus_grounding(),
        user_text,
        # 1024 truncated detailed multi-company answers mid-sentence (e.g. a risk
        # comparison table got cut off). 4096 comfortably fits the longest answers
        # we produce while staying well under Sonnet's output limit.
        max_tokens=4096,
        max_iters=MAX_TOOL_ITERS,
        on_text=on_text,
        on_tool_call=on_tool_call,
    )

    trace: list[dict] = [
        {
            "node": "agent",
            "type": "tool_call",
            "data": {"tool": tc.tool, "args": tc.args, "result": tc.result},
        }
        for tc in result.tool_calls
    ]
    usage = {"input_tokens": result.input_tokens, "output_tokens": result.output_tokens}
    answer = result.answer

    # Reliability floor: if the tool-loop yields no answer (e.g. a backend that
    # intermittently botches a tool call), fall back to plain synthesis over the
    # retrieved chunks β€” the proven /answer path, no tool-calling involved.
    if not answer.strip() and chunks:
        fb = synthesize(state["rewritten_query"], chunks)
        answer = fb.answer
        usage["input_tokens"] += fb.input_tokens
        usage["output_tokens"] += fb.output_tokens
        trace.append(
            {
                "node": "agent",
                "type": "fallback",
                "data": {"reason": "tool-loop produced no answer; synthesized from retrieved chunks"},
            }
        )

    trace.append({"node": "agent", "type": "synthesize", "data": {"answer": answer}})
    return {"answer": answer, "usage": usage, "trace": trace}