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
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:
@@ -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
@@ -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": [
|
||||
|
||||
@@ -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
@@ -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",
|
||||
},
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user