"""Provider-agnostic chat completion with native tool (function) calling. This module complements ``backend.ai`` (which handles the 16 stateless editor actions). It exposes a single ``chat_completion`` entry point used by the in-app agent loop: - OpenAI-compatible providers (DeepSeek, OpenRouter, NVIDIA, QwenCloud, Xiaomi, Mistral) use the ``tools`` / ``tool_calls`` protocol. - Google Gemini uses ``functionDeclarations`` / ``functionCall``. When a provider rejects the ``tools`` parameter (model without function calling support), the call is transparently retried without tools so callers degrade to plain chat (fallback protocol, see ``docs/AI_ARCHITECTURE_GUIDE.md``). """ from __future__ import annotations import json import logging import re from collections.abc import AsyncIterator from dataclasses import dataclass, field from typing import Any import httpx from backend.ai import PROVIDERS, _build_headers, _get_provider_config logger = logging.getLogger("obsigate.ai_chat") # Status codes that usually mean "tools not supported by this model". _TOOLS_UNSUPPORTED_STATUS = {400, 404, 422} @dataclass class ToolCall: """A single tool invocation requested by the model.""" id: str name: str arguments: dict[str, Any] = field(default_factory=dict) @dataclass class LLMResponse: """Normalized provider response (text and/or tool calls).""" content: str | None = None tool_calls: list[ToolCall] = field(default_factory=list) provider: str = "" model: str = "" @property def has_tool_calls(self) -> bool: return bool(self.tool_calls) def _parse_arguments(raw: Any) -> dict[str, Any]: """Parse tool-call arguments that may arrive as a JSON string or dict.""" if isinstance(raw, dict): return raw if isinstance(raw, str) and raw.strip(): try: parsed = json.loads(raw) return parsed if isinstance(parsed, dict) else {"value": parsed} except json.JSONDecodeError: return {"_raw": raw} return {} async def chat_completion( messages: list[dict[str, Any]], *, tools: list[dict[str, Any]] | None = None, provider: str | None = None, model: str | None = None, temperature: float = 0.3, max_tokens: int = 4096, ) -> LLMResponse: """Run a chat completion, optionally with tool calling. Args: messages: OpenAI-style messages (``system`` / ``user`` / ``assistant`` / ``tool``). Assistant messages may carry ``tool_calls``. tools: OpenAI-style tool schemas (see ``get_tool_schemas``). provider: Provider override; falls back to the configured default. model: Model override; falls back to the provider's default model. temperature: Sampling temperature. max_tokens: Maximum output tokens. Returns: :class:`LLMResponse` with the text content and/or requested tool calls. """ cfg = _get_provider_config(provider) # type: ignore[arg-type] name = cfg["name"] resolved_model = model or cfg["model"] if name == "gemini": return await _gemini_chat(messages, tools, resolved_model, temperature, max_tokens) return await _openai_chat(messages, tools, cfg, resolved_model, temperature, max_tokens) async def _openai_chat( messages: list[dict[str, Any]], tools: list[dict[str, Any]] | None, cfg: dict[str, Any], model: str, temperature: float, max_tokens: int, ) -> LLMResponse: """Call an OpenAI-compatible ``/chat/completions`` endpoint.""" headers = _build_headers(cfg) url = f"{cfg['base_url']}/chat/completions" def _build_payload(use_tools: list[dict[str, Any]] | None) -> dict[str, Any]: payload: dict[str, Any] = { "model": model, "messages": messages, "temperature": temperature, "max_tokens": max_tokens, } if use_tools: payload["tools"] = use_tools payload["tool_choice"] = "auto" return payload try: data = await _post_json(url, headers, _build_payload(tools)) except httpx.HTTPStatusError as e: if tools and e.response.status_code in _TOOLS_UNSUPPORTED_STATUS: logger.warning(f"Provider '{cfg['name']}' rejected tools — retrying without tools") data = await _post_json(url, headers, _build_payload(None)) else: raise message = data["choices"][0]["message"] content = (message.get("content") or "").strip() or None tool_calls: list[ToolCall] = [] for idx, tc in enumerate(message.get("tool_calls") or []): fn = tc.get("function", {}) or {} tool_calls.append(ToolCall( id=tc.get("id") or f"call_{idx}", name=fn.get("name", ""), arguments=_parse_arguments(fn.get("arguments")), )) return LLMResponse(content=content, tool_calls=tool_calls, provider=cfg["name"], model=model) def _gemini_tools(tools: list[dict[str, Any]] | None) -> list[dict[str, Any]] | None: """Convert OpenAI-style tool schemas to Gemini ``functionDeclarations``.""" if not tools: return None declarations = [] for spec in tools: fn = spec.get("function", spec) declarations.append({ "name": fn.get("name", ""), "description": fn.get("description", ""), "parameters": fn.get("parameters", {"type": "object", "properties": {}}), }) return [{"functionDeclarations": declarations}] _DATA_URL_RE = re.compile(r"^data:([^;,]+);base64,(.*)$", re.DOTALL) def _content_to_gemini_parts(content: Any) -> list[dict[str, Any]]: """Convert OpenAI-style message content to Gemini ``parts``. Accepts either a plain string or a multimodal content array (``[{"type": "text", ...}, {"type": "image_url", ...}]``). Data URLs are turned into ``inlineData`` parts so images can be sent to vision models. """ if content is None: return [{"text": ""}] if isinstance(content, str): return [{"text": content}] parts: list[dict[str, Any]] = [] if isinstance(content, list): for item in content: if isinstance(item, str): parts.append({"text": item}) continue if not isinstance(item, dict): continue item_type = item.get("type") if item_type == "text": parts.append({"text": item.get("text", "")}) elif item_type == "image_url": url = (item.get("image_url") or {}).get("url", "") match = _DATA_URL_RE.match(url or "") if match: parts.append({ "inlineData": {"mimeType": match.group(1), "data": match.group(2)}, }) elif url: parts.append({"fileData": {"fileUri": url}}) if not parts: parts.append({"text": ""}) return parts def _gemini_contents(messages: list[dict[str, Any]]) -> tuple[str, list[dict[str, Any]]]: """Split OpenAI-style messages into Gemini ``system`` text + ``contents``.""" system_parts: list[str] = [] contents: list[dict[str, Any]] = [] for msg in messages: role = msg.get("role") if role == "system": content = msg.get("content") or "" system_parts.append(content if isinstance(content, str) else "") elif role == "tool": contents.append({ "role": "user", "parts": [{ "functionResponse": { "name": msg.get("name", ""), "response": {"content": msg.get("content", "")}, }, }], }) else: contents.append({ "role": "user" if role == "user" else "model", "parts": _content_to_gemini_parts(msg.get("content")), }) return "\n".join(p for p in system_parts if p).strip(), contents async def _gemini_chat( messages: list[dict[str, Any]], tools: list[dict[str, Any]] | None, model: str, temperature: float, max_tokens: int, ) -> LLMResponse: """Call Gemini's ``generateContent`` endpoint with optional tools.""" cfg = PROVIDERS["gemini"] system, contents = _gemini_contents(messages) payload: dict[str, Any] = { "contents": contents, "generationConfig": {"temperature": temperature, "maxOutputTokens": max_tokens}, } if system: payload["system_instruction"] = {"parts": [{"text": system}]} gemini_tools = _gemini_tools(tools) if gemini_tools: payload["tools"] = gemini_tools url = f"{cfg['base_url']}/models/{model}:generateContent?key={cfg['api_key']}" data = await _post_json(url, None, payload) parts = data["candidates"][0]["content"].get("parts", []) content = "".join(p.get("text", "") for p in parts if "text" in p).strip() or None tool_calls: list[ToolCall] = [] for idx, part in enumerate(parts): fc = part.get("functionCall") if fc: tool_calls.append(ToolCall( id=f"call_{idx}", name=fc.get("name", ""), arguments=_parse_arguments(fc.get("args")), )) return LLMResponse(content=content, tool_calls=tool_calls, provider="gemini", model=model) async def _post_json(url: str, headers: dict[str, str] | None, payload: dict[str, Any]) -> dict[str, Any]: """POST JSON and return the parsed body, raising on HTTP errors.""" async with httpx.AsyncClient(timeout=120.0) as client: resp = await client.post(url, headers=headers or {}, json=payload) resp.raise_for_status() return resp.json() # ── Streaming (SSE token stream) ──────────────────────────────────────── async def stream_completion( messages: list[dict[str, Any]], *, provider: str | None = None, model: str | None = None, temperature: float = 0.3, max_tokens: int = 4096, ) -> AsyncIterator[str]: """Yield content deltas from a chat completion as they arrive. Only text content is streamed (no tool calling): this backs the plain ``/api/ai/bookslm/chat`` endpoint. The tool-calling ``/agent`` endpoint keeps using :func:`chat_completion` because tool calls need the complete response before they can be executed. """ cfg = _get_provider_config(provider) # type: ignore[arg-type] resolved_model = model or cfg["model"] if cfg["name"] == "gemini": stream = _gemini_stream(messages, resolved_model, temperature, max_tokens) else: stream = _openai_stream(messages, cfg, resolved_model, temperature, max_tokens) async for chunk in stream: yield chunk async def _openai_stream( messages: list[dict[str, Any]], cfg: dict[str, Any], model: str, temperature: float, max_tokens: int, ) -> AsyncIterator[str]: """Stream an OpenAI-compatible ``/chat/completions`` response.""" headers = _build_headers(cfg) url = f"{cfg['base_url']}/chat/completions" payload: dict[str, Any] = { "model": model, "messages": messages, "temperature": temperature, "max_tokens": max_tokens, "stream": True, } async with httpx.AsyncClient(timeout=120.0) as client, client.stream("POST", url, headers=headers, json=payload) as resp: resp.raise_for_status() async for line in resp.aiter_lines(): if not line or not line.startswith("data:"): continue raw = line[5:].strip() if raw == "[DONE]": break try: chunk = json.loads(raw) except json.JSONDecodeError: continue choices = chunk.get("choices") or [] if not choices: continue delta = choices[0].get("delta") or {} content = delta.get("content") if content: yield content async def _gemini_stream( messages: list[dict[str, Any]], model: str, temperature: float, max_tokens: int, ) -> AsyncIterator[str]: """Stream Gemini's ``streamGenerateContent`` response (SSE).""" cfg = PROVIDERS["gemini"] system, contents = _gemini_contents(messages) payload: dict[str, Any] = { "contents": contents, "generationConfig": {"temperature": temperature, "maxOutputTokens": max_tokens}, } if system: payload["system_instruction"] = {"parts": [{"text": system}]} url = f"{cfg['base_url']}/models/{model}:streamGenerateContent?alt=sse&key={cfg['api_key']}" async with httpx.AsyncClient(timeout=120.0) as client, client.stream("POST", url, json=payload) as resp: resp.raise_for_status() async for line in resp.aiter_lines(): if not line or not line.startswith("data:"): continue raw = line[5:].strip() if not raw: continue try: chunk = json.loads(raw) except json.JSONDecodeError: continue for candidate in chunk.get("candidates") or []: for part in candidate.get("content", {}).get("parts", []): text = part.get("text") if text: yield text