File size: 11,993 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
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
"""Gemini synthesis backend β€” the default provider.

Uses the unified `google-genai` SDK (NOT the legacy `google-generativeai`).
Same job as claude.py: retrieve β†’ synthesize a citation-grounded answer.
Maps Gemini's `usage_metadata` onto the shared `SynthesisResult`.

Why Gemini 2.5 Flash: the free tier zeroes out dev cost, and plain grounded
synthesis (read chunks, cite [N], don't hallucinate) is an easy workload for
it β€” this isn't reasoning-heavy. We disable "thinking" (budget=0) because
synthesis needs determinism and speed, not a scratchpad; thinking would just
burn output tokens and latency here.

Caching note: Gemini's implicit context cache only kicks in above a ~1k-token
prefix, and our system prompt is ~450 tokens, so cached_content_token_count
stays 0. That's expected, not a bug β€” see [[base]] SynthesisResult docstring.
"""

from __future__ import annotations

import re
import time
from functools import lru_cache

from google import genai
from google.genai import types

from finrag.config import settings
from finrag.llm.base import (
    MAX_TOKENS,
    SYSTEM_PROMPT,
    SynthesisResult,
    ToolCall,
    ToolLoopResult,
    build_user_message,
    empty_result,
    json_safe,
)
from finrag.retrieval.vector import RetrievedChunk

# JSON-schema lowercase types β†’ Gemini's uppercase Type enum values.
_TYPE_MAP = {
    "object": "OBJECT", "string": "STRING", "number": "NUMBER",
    "integer": "INTEGER", "boolean": "BOOLEAN", "array": "ARRAY",
}

# flash-lite is the default: the agent makes ~5 calls/question and 2.5-flash's
# free tier caps at only 20 requests/DAY, which an agentic workload exhausts in
# ~4 questions. flash-lite has a far larger free daily quota (~1000/day) and
# 15 req/min β€” enough to actually run and demo the agent for free. Quality is
# marginally lower but fine for grounded synthesis + mechanical sub-tasks.
# (3.5-flash resolves but 503s constantly on free tier; 2.5-flash is selectable
# by editing this line if billing is enabled.)
GEMINI_MODEL = "gemini-2.5-flash-lite"


@lru_cache(maxsize=1)
def get_gemini_client() -> genai.Client:
    if not settings.gemini_api_key:
        raise RuntimeError(
            "GEMINI_API_KEY is not set. Add it to .env, or set "
            "LLM_PROVIDER=anthropic to use Claude instead."
        )
    return genai.Client(api_key=settings.gemini_api_key)


def _retry_delay(exc: Exception, attempt: int) -> float | None:
    """Seconds to wait before retrying `exc`, or None if it's not retryable.

    Two transient free-tier failures:
      - 503 UNAVAILABLE ("high demand") β†’ linear backoff 2s/4s/6s.
      - 429 RESOURCE_EXHAUSTED (5 req/min cap) β†’ honor the API's suggested
        retryDelay (it tells us exactly when the per-minute window resets),
        with a small buffer and a sane ceiling.
    """
    s = str(exc)
    if "429" in s or "RESOURCE_EXHAUSTED" in s:
        # Per-DAY quota won't reset within any sane wait β€” fail fast so the
        # caller gets a clear error instead of blocking ~60s for nothing.
        # Only the per-minute cap is worth waiting out.
        if "PerDay" in s or "RequestsPerDay" in s:
            return None
        m = re.search(r"retry in ([0-9.]+)s", s) or re.search(
            r"retryDelay['\"]?:?\s*['\"]?([0-9.]+)s", s
        )
        return min((float(m.group(1)) if m else 20.0) + 1.0, 35.0)
    if "503" in s or "UNAVAILABLE" in s:
        return 2.0 * (attempt + 1)
    return None


def _has_content(response: object) -> bool:
    """True if the response carries at least one usable text/function_call part.

    flash-lite intermittently returns a candidate with no parts (empty
    response), especially on larger tool-laden prompts. Such a response isn't
    an exception, so we detect it explicitly and retry.
    """
    cands = getattr(response, "candidates", None)
    if not cands:
        return False
    cand = cands[0]
    if not cand.content or not cand.content.parts:
        return False
    return any(
        getattr(p, "text", None) or getattr(p, "function_call", None)
        for p in cand.content.parts
    )


def generate_content_with_retry(
    contents: object,
    config: types.GenerateContentConfig,
    *,
    retries: int = 5,
):
    """Single choke-point for Gemini calls, with retry on 503, per-minute 429,
    and empty (zero-part) responses.

    The agent (Decision 16) makes several calls per question; on the free tier
    any one can hit a transient 503, the per-minute 429, or a flash-lite empty
    candidate. Centralizing retry here means synthesis, NL→SQL, planning, and
    the tool-loop all inherit it, so a single blip doesn't abort the graph.
    """
    last: Exception | None = None
    for i in range(retries):
        try:
            response = get_gemini_client().models.generate_content(
                model=GEMINI_MODEL, contents=contents, config=config
            )
        except Exception as e:  # noqa: BLE001 β€” re-raised unless _retry_delay matches
            delay = _retry_delay(e, i)
            if delay is not None and i < retries - 1:
                last = e
                time.sleep(delay)
                continue
            raise
        # Empty candidate β†’ transient; retry a couple times before giving up.
        if not _has_content(response) and i < retries - 1:
            time.sleep(1.0)
            continue
        return response
    raise last  # type: ignore[misc]


def _config(
    system_instruction: str,
    *,
    tools: list[types.Tool] | None = None,
    max_output_tokens: int = MAX_TOKENS,
    temperature: float = 0.0,
) -> types.GenerateContentConfig:
    """Shared config: thinking disabled (synthesis/routing want determinism)."""
    return types.GenerateContentConfig(
        system_instruction=system_instruction,
        tools=tools,
        max_output_tokens=max_output_tokens,
        temperature=temperature,
        thinking_config=types.ThinkingConfig(thinking_budget=0),
    )


def _extract_text(response: object) -> tuple[str, str]:
    """Pull (text, finish_reason) out of a Gemini response defensively.

    If the candidate was blocked (safety) or truncated, `.parts` may be empty;
    `response.text` would warn/raise in that case, so we walk the parts.
    """
    text = ""
    finish_reason = "unknown"
    candidates = getattr(response, "candidates", None)
    if candidates:
        cand = candidates[0]
        finish_reason = str(getattr(cand, "finish_reason", "unknown"))
        if cand.content and cand.content.parts:
            text = "".join(
                p.text for p in cand.content.parts if getattr(p, "text", None)
            )
    return text, finish_reason


def generate_text(
    system_instruction: str,
    user_text: str,
    *,
    max_output_tokens: int = 512,
    temperature: float = 0.0,
) -> str:
    """Single-shot text completion β€” the building block for sub-LLM tasks
    like NL→SQL, where we want raw text out, not a SynthesisResult.

    (Lives on the Gemini backend for now; if LLM_PROVIDER swaps to Anthropic,
    this is the one helper sql_query would need mirrored in claude.py.)
    """
    response = generate_content_with_retry(
        user_text,
        _config(
            system_instruction,
            max_output_tokens=max_output_tokens,
            temperature=temperature,
        ),
    )
    text, _ = _extract_text(response)
    return text


def synthesize_gemini(question: str, chunks: list[RetrievedChunk]) -> SynthesisResult:
    if not chunks:
        return empty_result(GEMINI_MODEL)

    # system_instruction plays the role Anthropic's `system=` does β€” keeps
    # grounding rules out of the user turn. temperature 0 β†’ deterministic
    # grounded extraction, not creativity.
    response = generate_content_with_retry(
        build_user_message(question, chunks),
        _config(SYSTEM_PROMPT),
    )

    answer_text, finish_reason = _extract_text(response)

    usage = response.usage_metadata
    return SynthesisResult(
        answer=answer_text,
        model=GEMINI_MODEL,
        input_tokens=getattr(usage, "prompt_token_count", 0) or 0,
        output_tokens=getattr(usage, "candidates_token_count", 0) or 0,
        # Gemini doesn't bill a separate cache-write tier the way Anthropic
        # does; implicit caching just reports read tokens. Keep write at 0.
        cache_creation_input_tokens=0,
        cache_read_input_tokens=getattr(usage, "cached_content_token_count", 0) or 0,
        stop_reason=finish_reason,
    )


# ── Agent tool-loop (native function calling) ─────────────────────────────
def _to_schema(js: dict) -> types.Schema:
    """One JSON-schema fragment β†’ a genai Schema (recursively)."""
    schema = types.Schema(type=_TYPE_MAP.get(js.get("type", "object"), "STRING"))
    if "description" in js:
        schema.description = js["description"]
    if js.get("type") == "object":
        schema.properties = {k: _to_schema(v) for k, v in js.get("properties", {}).items()}
        if js.get("required"):
            schema.required = list(js["required"])
    if js.get("type") == "array" and "items" in js:
        schema.items = _to_schema(js["items"])
    return schema


def _gemini_tool() -> types.Tool:
    from finrag.tools import TOOL_SPECS  # lazy: avoid llm↔tools import cycle

    return types.Tool(
        function_declarations=[
            types.FunctionDeclaration(
                name=s.name, description=s.description, parameters=_to_schema(s.parameters)
            )
            for s in TOOL_SPECS
        ]
    )


def _args_to_dict(args: object) -> dict:
    """Convert a Gemini function_call.args (proto Map) to a plain dict."""
    def conv(v):
        if hasattr(v, "items"):
            return {k: conv(x) for k, x in v.items()}
        if isinstance(v, (list, tuple)):
            return [conv(x) for x in v]
        return v

    return conv(args) if args else {}


def tool_loop(
    system: str,
    user_text: str,
    *,
    max_tokens: int = 1024,
    max_iters: int = 5,
) -> ToolLoopResult:
    """Run Gemini with tools until it stops requesting them (or max_iters).
    Mirrors claude.tool_loop's signature/return so the dispatcher can pick."""
    from finrag.tools import dispatch  # lazy: avoid llm↔tools import cycle

    config = types.GenerateContentConfig(
        system_instruction=system,
        tools=[_gemini_tool()],
        max_output_tokens=max_tokens,
        temperature=0.0,
        thinking_config=types.ThinkingConfig(thinking_budget=0),
    )
    contents: list[types.Content] = [
        types.Content(role="user", parts=[types.Part(text=user_text)])
    ]
    in_tok = out_tok = 0
    calls: list[ToolCall] = []
    answer = ""

    for _ in range(max_iters):
        resp = generate_content_with_retry(contents, config)
        um = resp.usage_metadata
        in_tok += getattr(um, "prompt_token_count", 0) or 0
        out_tok += getattr(um, "candidates_token_count", 0) or 0

        cand = resp.candidates[0] if resp.candidates else None
        if cand is None or not cand.content or not cand.content.parts:
            break
        parts = cand.content.parts
        contents.append(cand.content)

        fcs = [p.function_call for p in parts if getattr(p, "function_call", None)]
        if not fcs:
            answer = "".join(p.text for p in parts if getattr(p, "text", None))
            break

        response_parts: list[types.Part] = []
        for fc in fcs:
            args = _args_to_dict(fc.args)
            result = json_safe(dispatch(fc.name, args))
            calls.append(ToolCall(fc.name, args, result))
            response_parts.append(
                types.Part.from_function_response(name=fc.name, response=result)
            )
        contents.append(types.Content(role="user", parts=response_parts))

    return ToolLoopResult(answer=answer, input_tokens=in_tok, output_tokens=out_tok, tool_calls=calls)