Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ The companion can be run as a local web app or a native MacOS app.

**Agent**
- Tight Text-to-SQL loop that uses query plans and schema stats to optimize queries.
- Anthropic and OpenAI providers with configurable prompts, step/token limits and tool timeouts
- Anthropic, OpenAI and Google providers with configurable prompts, step/token limits and tool timeouts
- Three chat modes:
- **SQL:** Turn a plain text request into a performant SQL query
- **Question:** Answer questions about your data in plain text
Expand Down
250 changes: 250 additions & 0 deletions backend/agent/gemini_provider.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,250 @@
"""Gemini-backed agent: NL request -> performant SQL via a tool-use loop.

Same contract as :class:`~backend.agent.anthropic_provider.AnthropicProvider`
(the :class:`~backend.agent.provider.AgentProvider` protocol), so the chat UI is
unchanged. Only the transport differs: the loop runs against the Gemini API, the
shared tool schemas become ``FunctionDeclaration``s, and the schema context goes
in ``system_instruction`` (Gemini caches long prefixes implicitly, so there are
no explicit cache markers).

The assistant turn is stored as the SDK's own ``Content`` object rather than a
rebuilt copy: Gemini 3 returns thought signatures alongside function calls and
rejects a follow-up turn that drops them.
"""

from __future__ import annotations

import threading
from typing import TYPE_CHECKING, Any, Iterator

from google import genai
from google.genai import errors as genai_errors
from google.genai import types

from backend.agent import context as ctx
from backend.agent.prompts import build_system_prompt
from backend.agent.provider import AgentEvent
from backend.agent.tools import execute_tool, tools_for_mode
from backend.config import DEFAULT_AGENT_MAX_STEPS, DEFAULT_AGENT_MAX_TOKENS
from backend.db.dialects import get_dialect

if TYPE_CHECKING:
from backend.agent.tools import SavedQueriesLoader
from backend.db.dialects.base import Dialect
from backend.db.extensions import ExtensionPlugin


def _to_gemini_tools(defs: list[dict]) -> list[types.Tool]:
"""Translate the shared Anthropic-style tool defs to Gemini function decls."""
declarations = [
types.FunctionDeclaration(
name=d["name"],
description=d["description"],
parameters_json_schema=d["input_schema"],
)
for d in defs
]
return [types.Tool(function_declarations=declarations)]


def _function_calls(content: types.Content | None) -> list[types.FunctionCall]:
parts = (content.parts if content else None) or []
return [p.function_call for p in parts if p.function_call]


class GeminiProvider:
def __init__(
self,
api_key: str,
model: str,
pool: Any,
plugins: "list[ExtensionPlugin] | None" = None,
dialect: "Dialect | None" = None,
custom_instructions: str = "",
statement_timeout_ms: int | None = None,
mode: str = "sql",
max_steps: int = DEFAULT_AGENT_MAX_STEPS,
max_tokens: int = DEFAULT_AGENT_MAX_TOKENS,
saved_queries_loader: "SavedQueriesLoader | None" = None,
):
self.client = genai.Client(api_key=api_key)
self.model = model
self.pool = pool
self.plugins = plugins or []
self.dialect = dialect or get_dialect(None)
self.custom_instructions = custom_instructions
self.statement_timeout_ms = statement_timeout_ms
self.mode = mode
self.max_steps = max_steps
self.max_tokens = max_tokens
self.saved_queries_loader = saved_queries_loader
# History holds Gemini Contents (roles "user"/"model"), without the
# system instruction — that is (re)built lazily and passed per request.
self.contents: list[types.Content] = []
self._system_text: str | None = None
# Set by stop() to interrupt an in-flight tool-use loop.
self._cancel = threading.Event()

def stop(self) -> None:
"""Request the current send() loop to halt at the next checkpoint."""
self._cancel.set()

def clear_stop(self) -> None:
"""Reset the cancel flag; call before starting a new send()."""
self._cancel.clear()

def reset(self) -> None:
self.contents = []
self._system_text = None

def seed_history(self, messages: list[dict]) -> None:
"""Replace history from persisted turns.

``messages`` is the alternating user/assistant *text* shape shared with
the other providers; Gemini names the assistant role "model".
"""
self.contents = [
types.Content(
role="model" if m["role"] == "assistant" else "user",
parts=[types.Part(text=m["content"])],
)
for m in messages
]

def refresh_context(self) -> None:
"""Drop the cached schema/stats context so it rebuilds next turn."""
self._system_text = None

def _system(self) -> str:
if self._system_text is None:
context = ctx.build_context(self.dialect, self.pool, plugins=self.plugins)
prompt = build_system_prompt(self.custom_instructions, self.mode)
self._system_text = f"{prompt}\n\n# Database context\n\n{context}"
return self._system_text

def _config(self) -> types.GenerateContentConfig:
return types.GenerateContentConfig(
system_instruction=self._system(),
max_output_tokens=self.max_tokens,
tools=_to_gemini_tools(tools_for_mode(self.mode)),
# We run the tool loop ourselves (the pool and timeouts live here),
# so the SDK must not call anything on our behalf.
automatic_function_calling=types.AutomaticFunctionCallingConfig(disable=True),
)

def _repair_history(self) -> None:
"""Ensure every model ``function_call`` is answered by a response part.

Gemini rejects a request where a function call has no matching
functionResponse. An interrupted send() (the SSE client disconnecting
mid-loop) can leave such orphans; we backfill synthetic error results so
the conversation stays valid.
"""
i = 0
while i < len(self.contents):
calls = _function_calls(self.contents[i]) if self.contents[i].role == "model" else []
if not calls:
i += 1
continue

# The responses to this turn, if any, are the immediately next Content.
covered: set[str] = set()
nxt = self.contents[i + 1] if i + 1 < len(self.contents) else None
if nxt is not None and nxt.role == "user":
covered = {
p.function_response.name
for p in (nxt.parts or [])
if p.function_response
}

orphaned = [c for c in calls if c.name not in covered]
if not orphaned:
i += 1
continue

synthetic = [
types.Part.from_function_response(
name=c.name, response={"error": "Interrupted."}
)
for c in orphaned
]
if nxt is not None and nxt.role == "user" and covered:
# Fold the missing responses into the partial answer turn.
nxt.parts = list(nxt.parts or []) + synthetic
else:
self.contents.insert(i + 1, types.Content(role="user", parts=synthetic))
i += 2

def send(self, user_message: str) -> Iterator[AgentEvent]:
self._repair_history()
self.contents.append(
types.Content(role="user", parts=[types.Part(text=user_message)])
)
try:
yield from self._run_loop()
except genai_errors.APIError as exc:
if exc.code in (401, 403):
yield AgentEvent("error", text="Invalid API key. Check Settings.")
elif exc.code == 429:
yield AgentEvent("error", text="Rate limited — wait a moment and retry.")
else:
yield AgentEvent("error", text=f"API error {exc.code}: {exc.message}")
except Exception as exc: # noqa: BLE001
yield AgentEvent("error", text=f"Agent error: {exc}")

def _run_loop(self) -> Iterator[AgentEvent]:
for _ in range(self.max_steps):
if self._cancel.is_set():
yield AgentEvent("done")
return
response = self.client.models.generate_content(
model=self.model,
contents=self.contents,
config=self._config(),
)
candidate = response.candidates[0] if response.candidates else None
content = candidate.content if candidate else None
parts = (content.parts if content else None) or []

# Keep the SDK's own Content: it carries the thought signatures Gemini
# requires back on the next turn alongside the function calls.
if content is not None:
self.contents.append(content)

for part in parts:
# Thought summaries are internal reasoning, not an answer.
if part.text and not part.thought and part.text.strip():
yield AgentEvent("text", text=part.text)

calls = _function_calls(content)
if not calls:
yield AgentEvent("done")
return

results: list[types.Part] = []
for call in calls:
# Stop before running further tools; unanswered calls are
# backfilled by _repair_history on the next send.
if self._cancel.is_set():
break
tool_input = dict(call.args or {})
yield AgentEvent("tool_call", tool_name=call.name, tool_input=tool_input)
result_text, is_error = execute_tool(
self.dialect, self.pool, call.name, tool_input,
statement_timeout_ms=self.statement_timeout_ms,
saved_queries_loader=self.saved_queries_loader,
)
yield AgentEvent("tool_result", tool_name=call.name, ok=not is_error)
results.append(
types.Part.from_function_response(
name=call.name,
response={"error" if is_error else "result": result_text},
)
)
self.contents.append(types.Content(role="user", parts=results))
if self._cancel.is_set():
yield AgentEvent("done")
return

yield AgentEvent("error", text="Stopped: too many tool iterations.")
yield AgentEvent("done")
21 changes: 18 additions & 3 deletions backend/agent/provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,20 +17,29 @@
from backend.db.dialects.base import Dialect
from backend.db.extensions import ExtensionPlugin

# Models offered in the UI, per provider. Anthropic defaults to Opus 4.8 — the
# EXPLAIN→evaluate→iterate loop benefits from its reasoning; the smaller models
# are cheaper, faster options. These are just the defaults surfaced in the
# Models offered in the UI, per provider, ordered most to least capable. The
# default (see DEFAULT_LLM_MODEL) is Haiku 4.5 — cheap and fast enough for most
# questions; the larger models are there for when the EXPLAIN→evaluate→iterate
# loop needs deeper reasoning. These are just the options surfaced in the
# picker; any model the provider accepts will run.
MODELS: dict[str, list[str]] = {
"anthropic": ["claude-fable-5", "claude-opus-4-8", "claude-sonnet-5", "claude-sonnet-4-6", "claude-haiku-4-5"],
# The GPT-5.6 family reasons by default, and /v1/chat/completions (see
# openai_provider) rejects function tools in that state — so they can't be
# offered until the provider moves to /v1/responses.
"openai": ["gpt-5", "gpt-5-mini", "gpt-4.1"],
"google": [
"gemini-3.1-pro-preview", "gemini-3.5-flash", "gemini-2.5-pro",
"gemini-2.5-flash", "gemini-2.5-flash-lite",
],
}

# Human-facing metadata for each provider (label + API-key placeholder), keyed
# in the same order they should appear in Settings.
PROVIDERS: dict[str, dict[str, str]] = {
"anthropic": {"label": "Anthropic", "keyPlaceholder": "sk-ant-…"},
"openai": {"label": "OpenAI", "keyPlaceholder": "sk-…"},
"google": {"label": "Google", "keyPlaceholder": "AIza…"},
}


Expand All @@ -45,6 +54,8 @@ def provider_for_model(model: str) -> str:
return "anthropic"
if model.startswith(("gpt", "o1", "o3", "o4")):
return "openai"
if model.startswith("gemini"):
return "google"
return next(iter(MODELS))


Expand Down Expand Up @@ -113,6 +124,10 @@ def build_provider(
from backend.agent.openai_provider import OpenAIProvider

cls = OpenAIProvider
elif provider == "google":
from backend.agent.gemini_provider import GeminiProvider

cls = GeminiProvider
else:
raise ValueError(f"Unknown provider: {provider}")

Expand Down
19 changes: 16 additions & 3 deletions frontend/src/components/chat/shared.ts
Original file line number Diff line number Diff line change
Expand Up @@ -161,12 +161,25 @@ export function oneLine(text: string, max = 200): string {
return s.length > max ? `${s.slice(0, max)}…` : s;
}

const title = (s: string) => s.charAt(0).toUpperCase() + s.slice(1);

// "claude-opus-4-8" -> "Opus 4.8", "claude-haiku-4-5-20251001" -> "Haiku 4.5".
// GPT ids keep their conventional casing: "gpt-5" -> "GPT-5", "gpt-4.1" -> "GPT-4.1".
// GPT and Gemini ids keep their conventional casing, with any trailing tier or
// codename spelled out: "gpt-4.1" -> "GPT-4.1", "gpt-5.6-sol" -> "GPT-5.6 Sol",
// "gemini-3.1-pro-preview" -> "Gemini 3.1 Pro Preview".
export function modelLabel(id: string): string {
if (/^(gpt|o\d)/i.test(id)) return id.replace(/^gpt/i, "GPT");
if (/^(gpt|o\d|gemini)/i.test(id)) {
const [family, ...rest] = id.split("-");
const version = /^\d/.test(rest[0] ?? "") ? rest.shift() : undefined;
const isGpt = /^gpt$/i.test(family);
const head = isGpt ? "GPT" : title(family);
// OpenAI hyphenates the version ("GPT-5"), Google spaces it ("Gemini 3.5").
const base = version ? `${head}${isGpt ? "-" : " "}${version}` : head;
const suffix = rest.map(title).join(" ");
return suffix ? `${base} ${suffix}` : base;
}
const parts = id.replace(/^claude-/, "").split("-");
const name = parts[0].charAt(0).toUpperCase() + parts[0].slice(1);
const name = title(parts[0]);
const nums = parts.slice(1).filter((p) => /^\d+$/.test(p) && p.length <= 2);
return nums.length ? `${name} ${nums.join(".")}` : name;
}
Expand Down
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ dependencies = [
"keyring>=25.0",
"pydantic>=2.0",
"openai>=1.40",
"google-genai>=1.0",
]

[project.optional-dependencies]
Expand Down
6 changes: 4 additions & 2 deletions tests/test_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,11 +39,13 @@ def test_non_loopback_host_is_forbidden(client):
def test_settings_reports_no_key_then_key(client):
body = client.get("/settings").json()
assert body["hasLlmKey"] is False
assert {p["id"] for p in body["providers"]} == {"anthropic", "openai"}
assert {p["id"] for p in body["providers"]} == {"anthropic", "openai", "google"}
assert all(p["hasKey"] is False for p in body["providers"])
assert "claude-opus-4-8" in body["models"]
assert "gpt-5" in body["models"]
assert "gemini-3.5-flash" in body["models"]
assert body["modelProviders"]["gpt-5"] == "openai"
assert body["modelProviders"]["gemini-3.5-flash"] == "google"
assert body["defaultModel"]

# Setting an OpenAI key flips only that provider's status.
Expand All @@ -53,7 +55,7 @@ def test_settings_reports_no_key_then_key(client):
body = client.get("/settings").json()
assert body["hasLlmKey"] is True
providers = {p["id"]: p["hasKey"] for p in body["providers"]}
assert providers == {"anthropic": False, "openai": True}
assert providers == {"anthropic": False, "openai": True, "google": False}


def test_agent_limits_default_roundtrip_and_clamp(client):
Expand Down
Loading
Loading