Files
ObsiGate/backend/bookslm_routes.py
T
bruno 4c4e415975
CI / lint (push) Successful in 58s
CI / security (push) Successful in 40s
CI / test (push) Successful in 1m15s
CI / build (push) Successful in 37s
CI / e2e (push) Successful in 10m15s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
feat(ai): services partages A2 + SSE streaming B4 + confirmations UI B5 (#79)
2026-09-11 17:06:40 -04:00

330 lines
12 KiB
Python

"""BooksLM API routes — directory-scoped AI chat for Obsidian vaults."""
import json
import logging
from pathlib import Path
from typing import Any
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, stream_completion
from backend.auth.middleware import check_vault_access, require_auth
from backend.bookslm import (
build_general_system_prompt,
build_system_prompt,
collect_directory_context,
collect_files_context,
empty_context,
)
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"])
VALID_MODES = {"directory", "documents", "general"}
# ── Request models ──
class BooksLMContextRequest(BaseModel):
vault: str | None = Field(default=None, description="Vault name (required for directory/documents modes)")
directory: str = Field(default="", description="Relative directory path within the vault")
mode: str = Field(default="directory", description="Context mode: 'directory', 'documents' or 'general'")
context_files: list[str] = Field(
default_factory=list,
description="Relative file paths to use as context (documents mode)",
)
class BooksLMChatRequest(BaseModel):
vault: str | None = Field(default=None, description="Vault name (required for directory/documents modes)")
directory: str = Field(default="", description="Relative directory path within the vault")
mode: str = Field(default="directory", description="Context mode: 'directory', 'documents' or 'general'")
context_files: list[str] = Field(
default_factory=list,
description="Relative file paths to use as context (documents mode)",
)
message: str = Field(description="User message")
conversation_history: list[dict[str, str]] = Field(
default_factory=list,
description="Previous conversation turns [{role, content}]",
)
provider: str | None = Field(
default=None,
description="AI provider override (e.g. 'deepseek', 'openrouter', 'gemini', 'nvidia', 'xiaomi', 'mistral', 'qwencloud'). "
"If not set, uses DEFAULT_PROVIDER.",
)
model: str | None = Field(
default=None,
description="Model name to use for this request. If not set, uses the provider's default model.",
)
confirm: dict[str, Any] | None = Field(
default=None,
description="Pending tool confirmation to apply (two-step propose/apply). "
"Shape: the ``error`` object of a previous ``confirmation`` event.",
)
confirm_messages: list[dict[str, Any]] | None = Field(
default=None,
description="Conversation snapshot returned alongside a ``confirmation`` event, "
"echoed back to resume the agent run.",
)
def _normalize_mode(mode: str | None) -> str:
mode = (mode or "directory").lower()
return mode if mode in VALID_MODES else "directory"
def _resolve_vault_path(vault: str | None, current_user):
"""Validate vault access and return (vault_name, vault_path)."""
if not vault:
raise HTTPException(status_code=400, detail="Champ 'vault' requis pour ce contexte")
if not check_vault_access(vault, current_user):
raise HTTPException(status_code=403, detail=f"Accès refusé à la vault '{vault}'")
vault_data = get_vault_data(vault)
if not vault_data:
raise HTTPException(status_code=404, detail=f"Vault '{vault}' not found")
return vault, Path(vault_data["path"])
def _build_context(mode: str, vault_path: Path | None, directory: str, context_files: list[str]):
"""Collect the context payload for the requested mode."""
if mode == "general":
return empty_context("general")
if mode == "documents":
ctx = collect_files_context(vault_path, context_files, scope="documents") # type: ignore[arg-type]
# No open document could be read → degrade gracefully to General.
return ctx if ctx["file_count"] else empty_context("general")
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":
_, vault_path = _resolve_vault_path(req.vault, current_user)
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 ──
@router.post("/context", response_model=BooksLMContextResponse)
async def api_bookslm_context(
req: BooksLMContextRequest,
current_user=Depends(require_auth),
):
"""Collect context for the AI assistant.
Three modes are supported:
* ``directory`` — every supported file under a vault directory;
* ``documents`` — only the explicitly listed open documents;
* ``general`` — no document context (assistant for the app itself).
"""
mode = _normalize_mode(req.mode)
if mode == "general":
return empty_context("general")
_, vault_path = _resolve_vault_path(req.vault, current_user)
return _build_context(mode, vault_path, req.directory, req.context_files)
@router.post(
"/chat",
response_class=StreamingResponse,
responses={200: {"content": {"text/event-stream": {}}, "description": "SSE token stream"}},
)
async def api_bookslm_chat(
req: BooksLMChatRequest,
current_user=Depends(require_auth),
):
"""Chat with the AI assistant.
Builds a system prompt from the selected context (directory, open
documents or general app knowledge), then streams the provider's answer
as Server-Sent Events.
"""
system_prompt = _resolve_system_prompt(req, current_user)
messages: list[dict[str, Any]] = [{"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})
async def generate_sse():
try:
# Resolve provider: explicit override wins, else first available.
cfg_name = _resolve_provider_name(req.provider)
if cfg_name is None:
err = "Aucun fournisseur AI configuré (clés API manquantes)"
error_data = json.dumps({"error": err}, ensure_ascii=False)
yield f"event: error\ndata: {error_data}\n\n"
return
# Stream token deltas as they arrive from the provider.
async for token in stream_completion(
messages,
provider=cfg_name,
model=req.model,
temperature=0.3,
max_tokens=4096,
):
data = json.dumps(
{"token": token, "provider": cfg_name, "model": req.model or ""},
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 chat 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",
},
)
@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) carrying the pending call
and the conversation snapshot; the client resumes by echoing them back in
``confirm`` / ``confirm_messages``.
"""
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,
resume_messages=req.confirm_messages,
confirm_pending=req.confirm,
)
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(
{"pending": result.pending or {}, "messages": result.messages},
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",
},
)