CI / lint (push) Successful in 57s
CI / security (push) Successful in 39s
CI / test (push) Successful in 1m13s
CI / build (push) Successful in 36s
CI / e2e (push) Successful in 10m13s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
- backend/ai_chat.py: chat_completion provider-agnostique (OpenAI-compat tools/tool_calls + Gemini functionDeclarations/functionCall), retry sans tools si rejete - backend/agent/loop.py: run_agent multi-etapes (limite 10, truncation, confirmation two-step), LLM injectable - endpoint opt-in POST /api/ai/bookslm/agent (events SSE tool/message/confirmation), extraction _resolve_system_prompt - tests: test_ai_chat.py, test_agent_loop.py + 3 tests endpoint (728 passed au total) - ROADMAP B1/B2/B3/B7 livres ; B4/B5/B6 restants
235 lines
8.0 KiB
Python
235 lines
8.0 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
|
|
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()
|