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
330 lines
12 KiB
Python
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",
|
|
},
|
|
)
|