feat(ai): phase B function calling in-app (agent loop + endpoint /agent)
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
This commit is contained in:
2026-09-11 12:45:16 -04:00
parent 400224a089
commit 55696bfb31
9 changed files with 1056 additions and 42 deletions
+187
View File
@@ -0,0 +1,187 @@
"""In-app agent loop — multi-step tool calling.
The loop drives an LLM that may request tool calls, executes them through the
shared tool layer (``backend.tools``), feeds the results back, and repeats
until the model produces a final answer or the iteration budget is exhausted.
The LLM is injected as an async callable so the loop is fully testable without
network access::
async def fake_llm(messages, tools):
return LLMResponse(content="done")
result = await run_agent(messages, ctx=ctx, llm=fake_llm)
"""
from __future__ import annotations
import json
import logging
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any
from backend.tools.api import (
ToolConfirmationRequired,
ToolContext,
ToolError,
ToolScope,
call_tool,
get_tool_schemas,
)
logger = logging.getLogger("obsigate.agent.loop")
DEFAULT_MAX_ITERATIONS = 10
# Cap the size of a tool result fed back to the model (chars).
MAX_TOOL_RESULT_CHARS = 100_000
# Stopping reasons
STOP_DONE = "done"
STOP_MAX_ITERATIONS = "max_iterations"
STOP_CONFIRMATION_REQUIRED = "confirmation_required"
@dataclass
class ToolCallRecord:
"""Audit-friendly record of one executed tool call."""
name: str
arguments: dict[str, Any]
ok: bool
result: Any
@dataclass
class AgentResult:
"""Outcome of an agent run."""
content: str = ""
messages: list[dict[str, Any]] = field(default_factory=list)
tool_calls: list[ToolCallRecord] = field(default_factory=list)
iterations: int = 0
stopped: str = STOP_DONE
pending: dict[str, Any] | None = None
async def provider_llm(messages: list[dict[str, Any]], tools: list[dict[str, Any]] | None):
"""Default LLM adapter backed by ``backend.ai_chat.chat_completion``."""
from backend.ai_chat import chat_completion
return await chat_completion(messages, tools=tools)
def _truncate(payload: Any) -> Any:
"""Truncate an oversized tool result before feeding it back to the model."""
serialized = json.dumps(payload, ensure_ascii=False, default=str)
if len(serialized) <= MAX_TOOL_RESULT_CHARS:
return payload
return {"truncated": True, "content": serialized[:MAX_TOOL_RESULT_CHARS]}
def _assistant_tool_message(content: str | None, tool_calls: list[Any]) -> dict[str, Any]:
"""Build the OpenAI-style assistant message carrying tool calls."""
return {
"role": "assistant",
"content": content or "",
"tool_calls": [
{
"id": call.id,
"type": "function",
"function": {
"name": call.name,
"arguments": json.dumps(call.arguments, ensure_ascii=False, default=str),
},
}
for call in tool_calls
],
}
async def run_agent(
messages: list[dict[str, Any]],
*,
ctx: ToolContext,
llm: Callable[..., Any] | None = None,
tools: list[dict[str, Any]] | None = None,
max_iterations: int = DEFAULT_MAX_ITERATIONS,
on_tool_call: Callable[[ToolCallRecord], None] | None = None,
) -> AgentResult:
"""Run the tool-calling loop until completion.
Args:
messages: Initial conversation (OpenAI-style), typically a system
message followed by the conversation history and the user message.
ctx: Tool execution context (identity, mode, confirmation state).
llm: Async callable ``(messages, tools) -> LLMResponse``. Defaults to
the real provider adapter.
tools: Tool schemas to expose. ``None`` exposes all in-app tools;
pass ``[]`` to disable tool calling (plain chat).
max_iterations: Hard cap on LLM round-trips.
on_tool_call: Optional callback invoked after each executed tool call.
Returns:
An :class:`AgentResult`. ``stopped`` is ``done``, ``max_iterations`` or
``confirmation_required`` (in which case ``pending`` holds the payload
to confirm, for the two-step propose/apply flow).
"""
llm = llm or provider_llm
if tools is None:
tools = get_tool_schemas(scope=ToolScope.IN_APP)
convo = [dict(m) for m in messages]
executed: list[ToolCallRecord] = []
for iteration in range(1, max_iterations + 1):
response = await llm(convo, tools)
if not response.has_tool_calls:
return AgentResult(
content=response.content or "",
messages=convo,
tool_calls=executed,
iterations=iteration,
stopped=STOP_DONE,
)
convo.append(_assistant_tool_message(response.content, response.tool_calls))
for call in response.tool_calls:
try:
result = call_tool(call.name, ctx, call.arguments)
payload = result.data
ok = True
except ToolConfirmationRequired as e:
logger.info(f"Agent paused: confirmation required for '{call.name}'")
return AgentResult(
content=response.content or "",
messages=convo,
tool_calls=executed,
iterations=iteration,
stopped=STOP_CONFIRMATION_REQUIRED,
pending=e.to_dict(),
)
except ToolError as e:
payload = e.to_dict()
ok = False
record = ToolCallRecord(name=call.name, arguments=call.arguments, ok=ok, result=payload)
executed.append(record)
if on_tool_call is not None:
on_tool_call(record)
convo.append({
"role": "tool",
"tool_call_id": call.id,
"name": call.name,
"content": json.dumps(_truncate(payload), ensure_ascii=False, default=str),
})
logger.warning(f"Agent reached max iterations ({max_iterations})")
return AgentResult(
content="",
messages=convo,
tool_calls=executed,
iterations=max_iterations,
stopped=STOP_MAX_ITERATIONS,
)
+19 -15
View File
@@ -117,7 +117,7 @@ DEFAULT_PROVIDER: ProviderName = os.getenv("AI_DEFAULT_PROVIDER", "deepseek") #
def get_default_provider() -> str:
"""Resolve the default provider: ``data/config.json`` > env > ``deepseek``."""
provider = _read_app_config().get("ai_default_provider") or os.getenv("AI_DEFAULT_PROVIDER", "deepseek")
provider: str = str(_read_app_config().get("ai_default_provider") or os.getenv("AI_DEFAULT_PROVIDER", "deepseek"))
return provider if provider in PROVIDERS else "deepseek"
@@ -152,6 +152,23 @@ def _get_provider_config(provider: ProviderName | None = None) -> dict:
return {"name": p, **cfg}
def _build_headers(cfg: dict) -> dict:
"""Build HTTP headers for an OpenAI-compatible provider config.
Most providers use ``Authorization: Bearer KEY``. Some (Xiaomi MiMo) use a
dedicated header like ``api-key: KEY`` — supported via the
``auth_header_name`` key in PROVIDERS (defaults to ``Authorization``).
"""
header_name = cfg.get("auth_header_name") or "Authorization"
header_value = cfg["auth_header"].format(api_key=cfg["api_key"])
if header_name == "Authorization" and not header_value.lower().startswith("bearer "):
header_value = "Bearer " + header_value
return {
header_name: header_value,
"Content-Type": "application/json",
}
async def _call_deepseek_openrouter(prompt: str, system: str, provider: ProviderName | None = None,
temperature: float = 0.7, max_tokens: int = 2048) -> str:
"""Call OpenAI-compatible API (DeepSeek, OpenRouter, Xiaomi MiMo, etc.)."""
@@ -159,20 +176,7 @@ async def _call_deepseek_openrouter(prompt: str, system: str, provider: Provider
# Debug: log masked key to diagnose 401
key_preview = cfg["api_key"][:8] + "..." + cfg["api_key"][-4:] if len(cfg["api_key"]) > 12 else "***"
logger.info(f"AI call: provider={cfg['name']} model={cfg['model']} key={key_preview}")
# Most providers use "Authorization: Bearer KEY". Some (Xiaomi MiMo) use a
# dedicated header like "api-key: KEY". We support both via the
# `auth_header_name` key in PROVIDERS — defaults to "Authorization".
header_name = cfg.get("auth_header_name") or "Authorization"
header_value = cfg["auth_header"].format(api_key=cfg["api_key"])
# If the auth_header template doesn't include "Bearer " but the default
# header is Authorization, prepend it. This preserves backward compatibility
# for providers that store just the raw key.
if header_name == "Authorization" and not header_value.lower().startswith("bearer "):
header_value = "Bearer " + header_value
headers = {
header_name: header_value,
"Content-Type": "application/json",
}
headers = _build_headers(cfg)
payload = {
"model": cfg["model"],
"messages": [
+234
View File
@@ -0,0 +1,234 @@
"""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()
+130 -20
View File
@@ -8,6 +8,8 @@ from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field
from backend.agent.loop import run_agent
from backend.ai_chat import chat_completion
from backend.auth.middleware import check_vault_access, require_auth
from backend.bookslm import (
build_general_system_prompt,
@@ -18,6 +20,7 @@ from backend.bookslm import (
)
from backend.indexer import get_vault_data, index
from backend.schemas import BooksLMContextResponse
from backend.tools.api import ToolContext, ToolMode
logger = logging.getLogger("obsigate.bookslm_routes")
router = APIRouter(prefix="/api/ai/bookslm", tags=["BooksLM"])
@@ -90,6 +93,42 @@ def _build_context(mode: str, vault_path: Path | None, directory: str, context_f
return collect_directory_context(vault_path, directory) # type: ignore[arg-type]
def _resolve_system_prompt(req, current_user) -> str:
"""Resolve the vault access and build the assistant system prompt.
Shared by the classic chat endpoint and the tool-calling agent endpoint.
"""
mode = _normalize_mode(req.mode)
vault_path: Path | None = None
if mode != "general":
_resolve_vault_path(req.vault, current_user)
vault_path = Path(get_vault_data(req.vault)["path"]) # type: ignore[index]
context = _build_context(mode, vault_path, req.directory, req.context_files)
effective_mode = context.get("scope", mode)
if effective_mode == "general":
return build_general_system_prompt(list(index.keys()))
if effective_mode == "documents":
return build_system_prompt(context, scope="documents")
if context["file_count"] == 0:
raise HTTPException(status_code=404, detail="Aucun fichier markdown trouvé dans ce dossier")
return build_system_prompt(context, scope="directory")
def _resolve_provider_name(requested: str | None) -> str | None:
"""Pick the provider to use: explicit override, else first available."""
from backend.ai import DEFAULT_PROVIDER, PROVIDERS
cfg_name = (requested or DEFAULT_PROVIDER).lower()
if cfg_name in PROVIDERS and PROVIDERS[cfg_name].get("api_key"):
return cfg_name
for pname, pcfg in PROVIDERS.items():
if pcfg.get("api_key") and pname != "gemini":
return pname
return None
# ── Endpoints ──
@@ -130,26 +169,7 @@ async def api_bookslm_chat(
documents or general app knowledge), then streams the provider's answer
as Server-Sent Events.
"""
mode = _normalize_mode(req.mode)
vault_path: Path | None = None
if mode != "general":
_resolve_vault_path(req.vault, current_user)
vault_path = Path(get_vault_data(req.vault)["path"]) # type: ignore[index]
# Collect context
context = _build_context(mode, vault_path, req.directory, req.context_files)
effective_mode = context.get("scope", mode)
# Build system prompt
if effective_mode == "general":
system_prompt = build_general_system_prompt(list(index.keys()))
elif effective_mode == "documents":
system_prompt = build_system_prompt(context, scope="documents")
else:
if context["file_count"] == 0:
raise HTTPException(status_code=404, detail="Aucun fichier markdown trouvé dans ce dossier")
system_prompt = build_system_prompt(context, scope="directory")
system_prompt = _resolve_system_prompt(req, current_user)
# Call AI provider
from backend.ai import DEFAULT_PROVIDER, PROVIDERS, _call_deepseek_openrouter, _call_gemini
@@ -226,3 +246,93 @@ async def api_bookslm_chat(
"X-Accel-Buffering": "no",
},
)
@router.post(
"/agent",
response_class=StreamingResponse,
responses={200: {"content": {"text/event-stream": {}}, "description": "SSE tool/agent stream"}},
)
async def api_bookslm_agent(
req: BooksLMChatRequest,
current_user=Depends(require_auth),
):
"""Chat with the tool-calling agent.
Same context as ``/chat`` but the model may call tools (read/search the
vault) through the shared tool layer. Emits one ``tool`` event per executed
tool call, then a final ``message`` event. Mutating tools pause the run with
a ``confirmation`` event (two-step propose/apply).
"""
system_prompt = _resolve_system_prompt(req, current_user)
messages: list[dict] = [{"role": "system", "content": system_prompt}]
for turn in req.conversation_history:
role = turn.get("role")
content = turn.get("content", "")
if role in ("user", "assistant") and content:
messages.append({"role": role, "content": content})
messages.append({"role": "user", "content": req.message})
ctx = ToolContext(user=current_user, mode=ToolMode.IN_APP)
async def _llm(msgs, tool_schemas):
return await chat_completion(
msgs,
tools=tool_schemas,
provider=req.provider,
model=req.model,
temperature=0.3,
max_tokens=4096,
)
async def generate_sse():
try:
cfg_name = _resolve_provider_name(req.provider)
if cfg_name is None:
error_data = json.dumps(
{"error": "Aucun fournisseur AI configuré (clés API manquantes)"},
ensure_ascii=False,
)
yield f"event: error\ndata: {error_data}\n\n"
return
result = await run_agent(messages, ctx=ctx, llm=_llm)
for rec in result.tool_calls:
payload = json.dumps(
{"name": rec.name, "ok": rec.ok, "arguments": rec.arguments},
ensure_ascii=False,
)
yield f"event: tool\ndata: {payload}\n\n"
if result.stopped == "confirmation_required":
pending = json.dumps(result.pending or {}, ensure_ascii=False)
yield f"event: confirmation\ndata: {pending}\n\n"
else:
data = json.dumps(
{
"token": result.content,
"provider": cfg_name,
"model": req.model or "",
"iterations": result.iterations,
"stopped": result.stopped,
},
ensure_ascii=False,
)
yield f"event: message\ndata: {data}\n\n"
yield "event: done\ndata: {}\n\n"
except Exception as e:
logger.error(f"BooksLM agent error: {e}")
error_data = json.dumps({"error": str(e)}, ensure_ascii=False)
yield f"event: error\ndata: {error_data}\n\n"
return StreamingResponse(
generate_sse(),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"Connection": "keep-alive",
"X-Accel-Buffering": "no",
},
)
+3 -1
View File
@@ -288,7 +288,9 @@ Voir `docs/ROADMAP.md` (item dédié) pour le détail des activités.
## 9. Références
- `backend/ai.py`, `backend/ai_routes.py` — couche fournisseurs + actions éditeur
- `backend/bookslm.py`, `backend/bookslm_routes.py` — assistant contextuel
- `backend/ai_chat.py` — chat completion provider-agnostique avec tool calling (OpenAI-compat + Gemini)
- `backend/agent/loop.py` — agent loop in-app (multi-étapes, LLM injectable)
- `backend/bookslm.py`, `backend/bookslm_routes.py` — assistant contextuel (+ endpoint `/agent`)
- `frontend/js/ai.js`, `frontend/js/bookslm.js` — UI IA
- `backend/auth/middleware.py` — permissions
- `backend/secret_redactor.py` — redaction
+6 -6
View File
@@ -699,14 +699,14 @@
- [x] **A4.** Audit : journalisation JSONL de chaque appel d'outil (qui, quoi, vault, résultat) — action `ai_tool_call`, arguments sensibles résumés
- [x] **A5.** Tests unitaires du registry + services (sans IA) — `tests/test_tools.py` (30 tests)
##### B. Function calling in-app (3-4 jours)
- [ ] **B1.** Abstraction tool-calling provider-agnostique : `_call_deepseek_openrouter` (`ai.py:115`) + `_call_gemini` (`ai.py:156`) + Ollama — payload `tools`/`tool_choice`, parsing `tool_calls` / `functionCall`
- [ ] **B2.** Agent loop `backend/agent/loop.py` : boucle tool→résultat→tool, limite d'itérations (10) + budget tokens
- [ ] **B3.** Fallback protocole texte `obsigate-action` si le modèle ne supporte pas les tools
- [ ] **B4.** SSE réellement streaming (corriger `bookslm_routes.py:211`)
##### B. Function calling in-app (3-4 jours) — 🔵 partiel (B1/B2/B3/B7 livrés 2026-09-11)
- [x] **B1.** Abstraction tool-calling provider-agnostique : `backend/ai_chat.py` (`chat_completion`, `ToolCall`, `LLMResponse`) — OpenAI-compat (`tools`/`tool_choice`, parsing `tool_calls`) + Gemini (`functionDeclarations`/`functionCall`)
- [x] **B2.** Agent loop `backend/agent/loop.py` : boucle tool→résultat→tool, limite d'itérations (10), truncation des résultats ; endpoint opt-in `POST /api/ai/bookslm/agent` (events SSE `tool`/`message`/`confirmation`)
- [x] **B3.** Fallback : retry sans `tools` si le provider rejette les tools (400/404/422) → chat simple ; protocole texte `obsigate-action` conservé côté frontend
- [ ] **B4.** SSE réellement streaming (corriger `bookslm_routes.py` — le message final reste envoyé en un seul événement)
- [ ] **B5.** Confirmations UI : outils `read` auto, outils `write` via carte Apply (`bookslm.js:515`), aperçu diff pour `edit_file`
- [ ] **B6.** Outils de navigation in-app : `open_file`, `reveal_in_tree` (événement `obsigate:open-file`)
- [ ] **B7.** Tests : agent loop (LLM mocké), confirmations, permissions
- [x] **B7.** Tests : agent loop LLM mocké (`tests/test_agent_loop.py`), providers (`tests/test_ai_chat.py`), endpoint (`tests/test_bookslm.py`)
##### C. Catalogue d'outils — lecture & recherche (1-2 jours)
- [ ] **C1.** Vaults/navigation : `list_vaults`, `list_directory`, `list_all_files`
+199
View File
@@ -0,0 +1,199 @@
# tests/test_agent_loop.py — Unit tests for the in-app agent loop (Phase B)
"""Tests for backend.agent.loop.run_agent using a scripted (mocked) LLM."""
import pytest
from backend.agent.loop import (
STOP_CONFIRMATION_REQUIRED,
STOP_DONE,
STOP_MAX_ITERATIONS,
run_agent,
)
from backend.ai_chat import LLMResponse, ToolCall
from backend.tools.api import ToolContext, ToolRisk
from backend.tools.registry import ToolSpec
from backend.tools.schemas import ListVaultsInput
def _ctx(vaults=None) -> ToolContext:
return ToolContext(
user={"username": "tester", "role": "admin", "vaults": vaults or ["*"]},
audit_enabled=False,
)
def _register(monkeypatch, name, handler, risk=ToolRisk.READ):
from backend.tools import registry
spec = ToolSpec(
name=name,
description=f"test tool {name}",
input_model=ListVaultsInput,
handler=handler,
risk=risk,
)
monkeypatch.setitem(registry._REGISTRY, name, spec)
return spec
class ScriptedLLM:
"""Async LLM returning pre-scripted responses and recording invocations."""
def __init__(self, responses):
self.responses = list(responses)
self.calls = []
async def __call__(self, messages, tools):
self.calls.append({"messages": [dict(m) for m in messages], "tools": tools})
return self.responses.pop(0)
# ═══════════════════════════════════════════════════════════════════
# Plain completion
# ═══════════════════════════════════════════════════════════════════
class TestPlainCompletion:
@pytest.mark.asyncio
async def test_no_tool_calls_returns_content(self):
llm = ScriptedLLM([LLMResponse(content="bonjour")])
result = await run_agent([{"role": "user", "content": "hi"}], ctx=_ctx(), llm=llm)
assert result.stopped == STOP_DONE
assert result.content == "bonjour"
assert result.iterations == 1
assert result.tool_calls == []
@pytest.mark.asyncio
async def test_tools_empty_disables_tool_calling(self):
llm = ScriptedLLM([LLMResponse(content="ok")])
await run_agent([{"role": "user", "content": "hi"}], ctx=_ctx(), llm=llm, tools=[])
assert llm.calls[0]["tools"] == []
@pytest.mark.asyncio
async def test_default_tools_expose_in_app_tools(self):
llm = ScriptedLLM([LLMResponse(content="ok")])
await run_agent([{"role": "user", "content": "hi"}], ctx=_ctx(), llm=llm)
names = {t["function"]["name"] for t in llm.calls[0]["tools"]}
assert "list_vaults" in names
# ═══════════════════════════════════════════════════════════════════
# Tool calling
# ═══════════════════════════════════════════════════════════════════
class TestToolCalling:
@pytest.mark.asyncio
async def test_executes_tool_then_answers(self, monkeypatch):
_register(monkeypatch, "_echo", lambda ctx, params: {"echoed": True})
llm = ScriptedLLM([
LLMResponse(tool_calls=[ToolCall(id="1", name="_echo", arguments={})]),
LLMResponse(content="final"),
])
result = await run_agent([{"role": "user", "content": "go"}], ctx=_ctx(), llm=llm)
assert result.stopped == STOP_DONE
assert result.content == "final"
assert result.iterations == 2
assert len(result.tool_calls) == 1
assert result.tool_calls[0].ok is True
assert result.tool_calls[0].result == {"echoed": True}
@pytest.mark.asyncio
async def test_tool_result_is_fed_back(self, monkeypatch):
_register(monkeypatch, "_echo", lambda ctx, params: {"value": 42})
llm = ScriptedLLM([
LLMResponse(tool_calls=[ToolCall(id="1", name="_echo", arguments={})]),
LLMResponse(content="done"),
])
await run_agent([{"role": "user", "content": "go"}], ctx=_ctx(), llm=llm)
# Second LLM call must contain a tool message with the result.
second_messages = llm.calls[1]["messages"]
tool_msgs = [m for m in second_messages if m.get("role") == "tool"]
assert len(tool_msgs) == 1
assert '"value": 42' in tool_msgs[0]["content"]
assert tool_msgs[0]["tool_call_id"] == "1"
@pytest.mark.asyncio
async def test_unknown_tool_recorded_as_error(self):
llm = ScriptedLLM([
LLMResponse(tool_calls=[ToolCall(id="1", name="does_not_exist", arguments={})]),
LLMResponse(content="recovered"),
])
result = await run_agent([{"role": "user", "content": "go"}], ctx=_ctx(), llm=llm)
assert result.stopped == STOP_DONE
assert len(result.tool_calls) == 1
assert result.tool_calls[0].ok is False
assert result.tool_calls[0].result["error"]["code"] == "not_found"
@pytest.mark.asyncio
async def test_on_tool_call_callback(self, monkeypatch):
_register(monkeypatch, "_echo", lambda ctx, params: "x")
seen = []
llm = ScriptedLLM([
LLMResponse(tool_calls=[ToolCall(id="1", name="_echo", arguments={})]),
LLMResponse(content="done"),
])
await run_agent([{"role": "user", "content": "go"}], ctx=_ctx(), llm=llm, on_tool_call=seen.append)
assert len(seen) == 1
assert seen[0].name == "_echo"
# ═══════════════════════════════════════════════════════════════════
# Confirmation & limits
# ═══════════════════════════════════════════════════════════════════
class TestConfirmationAndLimits:
@pytest.mark.asyncio
async def test_write_tool_pauses_for_confirmation(self, monkeypatch):
_register(monkeypatch, "_write", lambda ctx, params: {"done": True}, risk=ToolRisk.WRITE)
llm = ScriptedLLM([LLMResponse(tool_calls=[ToolCall(id="1", name="_write", arguments={})])])
result = await run_agent([{"role": "user", "content": "write"}], ctx=_ctx(), llm=llm)
assert result.stopped == STOP_CONFIRMATION_REQUIRED
assert result.pending is not None
assert result.pending["error"]["tool"] == "_write"
assert result.tool_calls == []
@pytest.mark.asyncio
async def test_write_tool_runs_when_confirmed(self, monkeypatch):
_register(monkeypatch, "_write", lambda ctx, params: {"done": True}, risk=ToolRisk.WRITE)
llm = ScriptedLLM([
LLMResponse(tool_calls=[ToolCall(id="1", name="_write", arguments={})]),
LLMResponse(content="done"),
])
ctx = _ctx()
ctx.confirmed = True
result = await run_agent([{"role": "user", "content": "write"}], ctx=ctx, llm=llm)
assert result.stopped == STOP_DONE
assert result.tool_calls[0].ok is True
@pytest.mark.asyncio
async def test_max_iterations_stops_loop(self, monkeypatch):
_register(monkeypatch, "_echo", lambda ctx, params: "x")
llm = ScriptedLLM([
LLMResponse(tool_calls=[ToolCall(id="1", name="_echo", arguments={})]),
LLMResponse(tool_calls=[ToolCall(id="2", name="_echo", arguments={})]),
LLMResponse(content="never reached"),
])
result = await run_agent([{"role": "user", "content": "loop"}], ctx=_ctx(), llm=llm, max_iterations=2)
assert result.stopped == STOP_MAX_ITERATIONS
assert result.iterations == 2
assert len(result.tool_calls) == 2
# ═══════════════════════════════════════════════════════════════════
# Permissions (index-backed)
# ═══════════════════════════════════════════════════════════════════
class TestAgentPermissions:
@pytest.mark.asyncio
async def test_permission_denied_recorded(self, client):
llm = ScriptedLLM([
LLMResponse(tool_calls=[ToolCall(id="1", name="list_directory", arguments={"vault": "TestVault"})]),
LLMResponse(content="denied"),
])
ctx = ToolContext(user={"username": "limited", "vaults": ["OtherVault"]}, audit_enabled=False)
result = await run_agent([{"role": "user", "content": "list"}], ctx=ctx, llm=llm)
assert result.tool_calls[0].ok is False
assert result.tool_calls[0].result["error"]["code"] == "vault_access_denied"
+201
View File
@@ -0,0 +1,201 @@
# tests/test_ai_chat.py — Unit tests for provider-agnostic tool calling (Phase B)
"""Tests for backend.ai_chat: payload building, parsing, and tools fallback."""
import httpx
import pytest
from backend import ai_chat
from backend.ai_chat import ToolCall, _gemini_tools, _parse_arguments, chat_completion
OPENAI_CFG = {
"name": "deepseek",
"api_key": "sk-test-key",
"base_url": "https://api.deepseek.com/v1",
"model": "deepseek-chat",
"auth_header": "Bearer {api_key}",
}
GEMINI_CFG = {
"name": "gemini",
"api_key": "gem-test",
"base_url": "https://generativelanguage.googleapis.com/v1beta",
"model": "gemini-2.0-flash",
"auth_header": None,
}
TOOLS = [{
"type": "function",
"function": {
"name": "read_file",
"description": "Read a file",
"parameters": {"type": "object", "properties": {"vault": {"type": "string"}}, "required": ["vault"]},
},
}]
class _Capture:
"""Fake ``_post_json`` capturing calls and returning scripted responses."""
def __init__(self, responses):
self.responses = list(responses)
self.calls = []
async def __call__(self, url, headers, payload):
self.calls.append({"url": url, "headers": headers, "payload": payload})
result = self.responses.pop(0)
if isinstance(result, Exception):
raise result
return result
def _http_error(status: int) -> httpx.HTTPStatusError:
request = httpx.Request("POST", "https://example.test")
response = httpx.Response(status, request=request)
return httpx.HTTPStatusError("error", request=request, response=response)
# ═══════════════════════════════════════════════════════════════════
# Argument parsing
# ═══════════════════════════════════════════════════════════════════
class TestParseArguments:
def test_dict_passthrough(self):
assert _parse_arguments({"a": 1}) == {"a": 1}
def test_json_string(self):
assert _parse_arguments('{"a": 1}') == {"a": 1}
def test_invalid_string(self):
assert _parse_arguments("not json") == {"_raw": "not json"}
def test_empty_string(self):
assert _parse_arguments("") == {}
def test_none(self):
assert _parse_arguments(None) == {}
# ═══════════════════════════════════════════════════════════════════
# OpenAI-compatible path
# ═══════════════════════════════════════════════════════════════════
class TestOpenAIChat:
@pytest.mark.asyncio
async def test_content_only(self, monkeypatch):
monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG)
cap = _Capture([{"choices": [{"message": {"content": "bonjour"}}]}])
monkeypatch.setattr(ai_chat, "_post_json", cap)
result = await chat_completion([{"role": "user", "content": "hi"}])
assert result.content == "bonjour"
assert result.tool_calls == []
assert result.provider == "deepseek"
@pytest.mark.asyncio
async def test_tool_calls_parsed(self, monkeypatch):
monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG)
cap = _Capture([{
"choices": [{
"message": {
"content": None,
"tool_calls": [{
"id": "call_1",
"type": "function",
"function": {"name": "read_file", "arguments": '{"vault": "V", "path": "a.md"}'},
}],
},
}],
}])
monkeypatch.setattr(ai_chat, "_post_json", cap)
result = await chat_completion([{"role": "user", "content": "hi"}])
assert result.content is None
assert result.has_tool_calls
assert result.tool_calls[0] == ToolCall(id="call_1", name="read_file", arguments={"vault": "V", "path": "a.md"})
@pytest.mark.asyncio
async def test_tools_included_in_payload(self, monkeypatch):
monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG)
cap = _Capture([{"choices": [{"message": {"content": "ok"}}]}])
monkeypatch.setattr(ai_chat, "_post_json", cap)
await chat_completion([{"role": "user", "content": "hi"}], tools=TOOLS)
payload = cap.calls[0]["payload"]
assert payload["tools"] == TOOLS
assert payload["tool_choice"] == "auto"
@pytest.mark.asyncio
async def test_fallback_when_tools_rejected(self, monkeypatch):
monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG)
cap = _Capture([
_http_error(400),
{"choices": [{"message": {"content": "sans outils"}}]},
])
monkeypatch.setattr(ai_chat, "_post_json", cap)
result = await chat_completion([{"role": "user", "content": "hi"}], tools=TOOLS)
assert result.content == "sans outils"
assert "tools" in cap.calls[0]["payload"]
assert "tools" not in cap.calls[1]["payload"]
@pytest.mark.asyncio
async def test_http_error_without_tools_propagates(self, monkeypatch):
monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: OPENAI_CFG)
cap = _Capture([_http_error(500)])
monkeypatch.setattr(ai_chat, "_post_json", cap)
with pytest.raises(httpx.HTTPStatusError):
await chat_completion([{"role": "user", "content": "hi"}])
# ═══════════════════════════════════════════════════════════════════
# Gemini path
# ═══════════════════════════════════════════════════════════════════
class TestGeminiChat:
@pytest.mark.asyncio
async def test_content_and_system_instruction(self, monkeypatch):
monkeypatch.setitem(ai_chat.PROVIDERS, "gemini", GEMINI_CFG)
monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: GEMINI_CFG)
cap = _Capture([{"candidates": [{"content": {"parts": [{"text": "salut"}]}}]}])
monkeypatch.setattr(ai_chat, "_post_json", cap)
result = await chat_completion([
{"role": "system", "content": "sys"},
{"role": "user", "content": "hi"},
])
assert result.content == "salut"
assert result.provider == "gemini"
payload = cap.calls[0]["payload"]
assert payload["system_instruction"]["parts"][0]["text"] == "sys"
@pytest.mark.asyncio
async def test_function_call_parsed(self, monkeypatch):
monkeypatch.setitem(ai_chat.PROVIDERS, "gemini", GEMINI_CFG)
monkeypatch.setattr(ai_chat, "_get_provider_config", lambda provider=None: GEMINI_CFG)
cap = _Capture([{
"candidates": [{"content": {"parts": [
{"functionCall": {"name": "read_file", "args": {"vault": "V", "path": "a.md"}}},
]}}],
}])
monkeypatch.setattr(ai_chat, "_post_json", cap)
result = await chat_completion([{"role": "user", "content": "hi"}], tools=TOOLS)
assert result.tool_calls[0].name == "read_file"
assert result.tool_calls[0].arguments == {"vault": "V", "path": "a.md"}
# Tools converted to Gemini declarations
decls = cap.calls[0]["payload"]["tools"][0]["functionDeclarations"]
assert decls[0]["name"] == "read_file"
class TestGeminiToolsConversion:
def test_none(self):
assert _gemini_tools(None) is None
def test_conversion(self):
converted = _gemini_tools(TOOLS)
assert converted[0]["functionDeclarations"][0]["name"] == "read_file"
assert converted[0]["functionDeclarations"][0]["description"] == "Read a file"
+77
View File
@@ -746,3 +746,80 @@ class TestBooksLMChatModes:
)
assert resp.status_code in (200, 500), f"unexpected {resp.status_code}: {resp.text[:200]}"
assert resp.status_code != 422
# ── Tool-calling agent endpoint ───────────────────────────────────────
class TestBooksLMAgentEndpoint:
"""Tests for POST /api/ai/bookslm/agent (Phase B)."""
def test_agent_requires_auth(self, bookslm_client):
resp = bookslm_client.post(
"/api/ai/bookslm/agent",
json={"vault": "TestVault", "directory": "", "message": "Hello"},
)
assert resp.status_code == 401
def test_agent_tool_call_flow(self, bookslm_client, monkeypatch):
import backend.bookslm_routes as routes
from backend.ai_chat import LLMResponse, ToolCall
responses = [
LLMResponse(tool_calls=[ToolCall(id="1", name="list_directory", arguments={"vault": "TestVault"})]),
LLMResponse(content="Il y a des fichiers."),
]
async def fake_chat_completion(messages, **kwargs):
return responses.pop(0)
monkeypatch.setattr(routes, "chat_completion", fake_chat_completion)
monkeypatch.setattr(routes, "_resolve_provider_name", lambda requested: "deepseek")
token, _ = _login_bookslm(bookslm_client)
resp = bookslm_client.post(
"/api/ai/bookslm/agent",
json={"vault": "TestVault", "directory": "", "message": "liste les fichiers", "mode": "directory"},
headers={"Authorization": f"Bearer {token}"},
)
assert resp.status_code == 200
body = resp.text
assert "event: tool" in body
assert "list_directory" in body
assert "event: message" in body
assert "Il y a des fichiers." in body
def test_agent_pauses_for_confirmation(self, bookslm_client, monkeypatch):
import backend.bookslm_routes as routes
from backend.ai_chat import LLMResponse, ToolCall
from backend.tools import registry
from backend.tools.api import ToolRisk
from backend.tools.registry import ToolSpec
from backend.tools.schemas import ListVaultsInput
spec = ToolSpec(
name="_agent_write",
description="write for tests",
input_model=ListVaultsInput,
handler=lambda ctx, params: {"ok": True},
risk=ToolRisk.WRITE,
)
monkeypatch.setitem(registry._REGISTRY, "_agent_write", spec)
responses = [LLMResponse(tool_calls=[ToolCall(id="1", name="_agent_write", arguments={})])]
async def fake_chat_completion(messages, **kwargs):
return responses.pop(0)
monkeypatch.setattr(routes, "chat_completion", fake_chat_completion)
monkeypatch.setattr(routes, "_resolve_provider_name", lambda requested: "deepseek")
token, _ = _login_bookslm(bookslm_client)
resp = bookslm_client.post(
"/api/ai/bookslm/agent",
json={"vault": "TestVault", "directory": "", "message": "crée un fichier", "mode": "directory"},
headers={"Authorization": f"Bearer {token}"},
)
assert resp.status_code == 200
assert "event: confirmation" in resp.text
assert "_agent_write" in resp.text