"""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 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}] 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": system_parts.append(msg.get("content") or "") 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": [{"text": msg.get("content") or ""}], }) 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()