Files
bruno f049e208b6
CI / lint (push) Successful in 1m10s
CI / security (push) Successful in 43s
CI / test (push) Successful in 2m34s
CI / build (push) Successful in 43s
CI / e2e (push) Successful in 10m59s
feat(ai): commandes @/ & skills, analyse d'images et capacites des modeles (#81)
2026-09-12 11:38:46 -04:00

383 lines
13 KiB
Python

"""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