feat(ai): services partages A2 + SSE streaming B4 + confirmations UI B5 (#79)
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
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
This commit is contained in:
+11
-21
@@ -15,6 +15,8 @@ from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from backend.auth.middleware import check_vault_access
|
||||
from backend.services.errors import ServiceError
|
||||
from backend.services.paths import resolve_safe_path as _resolve_service_path
|
||||
|
||||
logger = logging.getLogger("obsigate.tools")
|
||||
|
||||
@@ -138,32 +140,20 @@ class ToolContext:
|
||||
)
|
||||
|
||||
|
||||
def resolve_safe_path(vault_root: Path, relative_path: str) -> Path:
|
||||
def resolve_safe_path(vault_root: Path, relative_path: str | None) -> Path:
|
||||
"""Resolve a vault-relative path, rejecting traversal outside the vault.
|
||||
|
||||
Mirrors ``backend.main._resolve_safe_path`` but raises a domain error
|
||||
instead of ``HTTPException`` so the tool layer stays transport-agnostic.
|
||||
Delegates to the shared :func:`backend.services.paths.resolve_safe_path`
|
||||
and maps its :class:`ServiceError` to tool domain errors so the tool layer
|
||||
stays transport-agnostic.
|
||||
|
||||
Raises:
|
||||
ToolPermissionError: When the resolved path escapes the vault root.
|
||||
ToolError: When the path cannot be resolved.
|
||||
"""
|
||||
full_path = vault_root / (relative_path or "")
|
||||
try:
|
||||
resolved = full_path.resolve(strict=False)
|
||||
root = vault_root.resolve(strict=False)
|
||||
except Exception as e:
|
||||
raise ToolError(f"Path resolution error: {e}", code="path_error") from e
|
||||
|
||||
try:
|
||||
resolved.relative_to(root)
|
||||
except ValueError:
|
||||
# Case-insensitive fallback for Windows / Docker path casing.
|
||||
if not str(resolved).lower().startswith(str(root).lower()):
|
||||
logger.warning(f"Path outside vault: vault={root}, requested={relative_path}")
|
||||
raise ToolPermissionError(
|
||||
"Access denied: path outside vault",
|
||||
code="path_outside_vault",
|
||||
details={"path": relative_path},
|
||||
) from None
|
||||
return resolved
|
||||
return _resolve_service_path(vault_root, relative_path)
|
||||
except ServiceError as e:
|
||||
if e.code in ("path_outside_vault", "permission_denied"):
|
||||
raise ToolPermissionError(e.message, code=e.code, details=e.details) from e
|
||||
raise ToolError(e.message, code=e.code, details=e.details) from e
|
||||
|
||||
@@ -18,6 +18,7 @@ from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from backend.services.errors import ServiceError
|
||||
from backend.tools.audit import log_tool_call
|
||||
from backend.tools.context import (
|
||||
ToolConfirmationRequired,
|
||||
@@ -136,6 +137,17 @@ def _audit(ctx: ToolContext, spec: ToolSpec, arguments: dict[str, Any], *, ok: b
|
||||
logger.debug(f"Tool audit failed for '{spec.name}': {e}")
|
||||
|
||||
|
||||
def _map_service_error(e: ServiceError) -> ToolError:
|
||||
"""Map a shared-layer :class:`ServiceError` to a tool domain error."""
|
||||
if e.code == "not_found":
|
||||
return ToolNotFoundError(e.message, details=e.details)
|
||||
if e.code in ("permission_denied", "path_outside_vault", "vault_access_denied"):
|
||||
return ToolPermissionError(e.message, code=e.code, details=e.details)
|
||||
if e.code == "invalid_arguments":
|
||||
return ToolValidationError(e.message, details=e.details)
|
||||
return ToolError(e.message, code=e.code, details=e.details)
|
||||
|
||||
|
||||
def call_tool(
|
||||
name: str,
|
||||
ctx: ToolContext,
|
||||
@@ -191,6 +203,10 @@ def call_tool(
|
||||
except ToolError as e:
|
||||
_audit(ctx, spec, arguments, ok=False, error=e.code)
|
||||
raise
|
||||
except ServiceError as e:
|
||||
mapped = _map_service_error(e)
|
||||
_audit(ctx, spec, arguments, ok=False, error=mapped.code)
|
||||
raise mapped from e
|
||||
except Exception as e:
|
||||
logger.error(f"Tool '{name}' failed: {e}")
|
||||
exec_error = ToolError(f"Tool '{name}' failed: {e}", code="tool_execution_error")
|
||||
|
||||
+26
-79
@@ -1,24 +1,21 @@
|
||||
"""Built-in tool services (Phase 0 — read/search).
|
||||
|
||||
These functions are the single source of truth consumed by both the in-app
|
||||
assistant and the MCP server. They delegate to the existing core modules
|
||||
(``backend.indexer``, ``backend.search``) rather than duplicating route logic.
|
||||
assistant and the MCP server. They delegate to the shared business-logic
|
||||
services (``backend.services``) so routes and tools never diverge.
|
||||
Mutating tools (create/edit/rename/move/delete) are added in later phases.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from backend.tools.context import (
|
||||
ToolContext,
|
||||
ToolError,
|
||||
ToolNotFoundError,
|
||||
ToolRisk,
|
||||
resolve_safe_path,
|
||||
)
|
||||
from backend.services.files import read_file_text
|
||||
from backend.services.search import list_tags as _list_tags
|
||||
from backend.services.search import search_vaults
|
||||
from backend.services.vaults import browse_directory, list_accessible_vaults
|
||||
from backend.tools.context import ToolContext, ToolRisk
|
||||
from backend.tools.registry import tool
|
||||
from backend.tools.schemas import (
|
||||
ListDirectoryInput,
|
||||
@@ -34,15 +31,6 @@ logger = logging.getLogger("obsigate.tools.service")
|
||||
TOOL_MAX_READ_BYTES = 200_000
|
||||
|
||||
|
||||
def _vault_data_or_raise(vault: str) -> dict[str, Any]:
|
||||
from backend.indexer import get_vault_data
|
||||
|
||||
data = get_vault_data(vault)
|
||||
if not data:
|
||||
raise ToolNotFoundError(f"Vault '{vault}' not found", details={"vault": vault})
|
||||
return data
|
||||
|
||||
|
||||
@tool(
|
||||
name="list_vaults",
|
||||
description="List the vaults the current user is allowed to access.",
|
||||
@@ -51,15 +39,10 @@ def _vault_data_or_raise(vault: str) -> dict[str, Any]:
|
||||
)
|
||||
def list_vaults(ctx: ToolContext, _params: ListVaultsInput) -> list[dict[str, Any]]:
|
||||
"""Return accessible vaults with a file count."""
|
||||
from backend.indexer import get_vault_data, get_vault_names
|
||||
|
||||
vaults: list[dict[str, Any]] = []
|
||||
for name in get_vault_names():
|
||||
if not ctx.has_vault_access(name):
|
||||
continue
|
||||
data = get_vault_data(name) or {}
|
||||
vaults.append({"name": name, "file_count": len(data.get("files", []))})
|
||||
return vaults
|
||||
return [
|
||||
{"name": v["name"], "file_count": v["file_count"]}
|
||||
for v in list_accessible_vaults(ctx.user)
|
||||
]
|
||||
|
||||
|
||||
@tool(
|
||||
@@ -71,28 +54,11 @@ def list_vaults(ctx: ToolContext, _params: ListVaultsInput) -> list[dict[str, An
|
||||
)
|
||||
def list_directory(ctx: ToolContext, params: ListDirectoryInput) -> list[dict[str, Any]]:
|
||||
"""Return the entries of a vault directory (direct children only)."""
|
||||
data = _vault_data_or_raise(params.vault)
|
||||
root = Path(data["path"])
|
||||
target = resolve_safe_path(root, params.path) if params.path else root.resolve()
|
||||
|
||||
if not target.exists() or not target.is_dir():
|
||||
raise ToolNotFoundError(f"Directory not found: {params.path}", details={"vault": params.vault, "path": params.path})
|
||||
|
||||
from backend.vault_settings import get_vault_setting
|
||||
|
||||
hide_hidden = (get_vault_setting(params.vault) or {}).get("hideHiddenFiles", False)
|
||||
|
||||
items: list[dict[str, Any]] = []
|
||||
for entry in sorted(target.iterdir(), key=lambda e: (not e.is_dir(), e.name.lower())):
|
||||
if hide_hidden and entry.name.startswith("."):
|
||||
continue
|
||||
rel = str(entry.relative_to(root)).replace("\\", "/")
|
||||
items.append({
|
||||
"name": entry.name,
|
||||
"path": rel,
|
||||
"type": "directory" if entry.is_dir() else "file",
|
||||
})
|
||||
return items
|
||||
data = browse_directory(params.vault, params.path)
|
||||
return [
|
||||
{"name": item["name"], "path": item["path"], "type": item["type"]}
|
||||
for item in data["items"]
|
||||
]
|
||||
|
||||
|
||||
@tool(
|
||||
@@ -104,27 +70,12 @@ def list_directory(ctx: ToolContext, params: ListDirectoryInput) -> list[dict[st
|
||||
)
|
||||
def read_file(ctx: ToolContext, params: ReadFileInput) -> dict[str, Any]:
|
||||
"""Return the (redacted) content of a vault file."""
|
||||
data = _vault_data_or_raise(params.vault)
|
||||
root = Path(data["path"])
|
||||
target = resolve_safe_path(root, params.path)
|
||||
|
||||
if not target.exists() or not target.is_file():
|
||||
raise ToolNotFoundError(f"File not found: {params.path}", details={"vault": params.vault, "path": params.path})
|
||||
|
||||
size = target.stat().st_size
|
||||
if size > TOOL_MAX_READ_BYTES:
|
||||
raise ToolError(
|
||||
f"File too large ({size} bytes > {TOOL_MAX_READ_BYTES})",
|
||||
code="file_too_large",
|
||||
details={"vault": params.vault, "path": params.path, "size": size},
|
||||
)
|
||||
|
||||
content = target.read_text(encoding="utf-8", errors="replace")
|
||||
|
||||
from backend.secret_redactor import redact_file_content
|
||||
|
||||
content = redact_file_content(content, params.path)
|
||||
return {"vault": params.vault, "path": params.path, "size": size, "content": content}
|
||||
return read_file_text(
|
||||
params.vault,
|
||||
params.path,
|
||||
redact=True,
|
||||
max_bytes=TOOL_MAX_READ_BYTES,
|
||||
)
|
||||
|
||||
|
||||
@tool(
|
||||
@@ -135,10 +86,8 @@ def read_file(ctx: ToolContext, params: ReadFileInput) -> dict[str, Any]:
|
||||
)
|
||||
def search_fulltext(ctx: ToolContext, params: SearchFulltextInput) -> list[dict[str, Any]]:
|
||||
"""Return ranked search results, filtered to accessible vaults."""
|
||||
from backend.search import search
|
||||
|
||||
results = search(params.q, vault_filter=params.vault, tag_filter=params.tag, limit=params.limit)
|
||||
return [r for r in results if ctx.has_vault_access(r["vault"])]
|
||||
payload = search_vaults(params.q, vault=params.vault, tag=params.tag, limit=params.limit)
|
||||
return [r for r in payload["results"] if ctx.has_vault_access(r["vault"])]
|
||||
|
||||
|
||||
@tool(
|
||||
@@ -149,11 +98,9 @@ def search_fulltext(ctx: ToolContext, params: SearchFulltextInput) -> list[dict[
|
||||
)
|
||||
def list_tags(ctx: ToolContext, params: ListTagsInput) -> list[dict[str, Any]]:
|
||||
"""Return tags sorted by descending count."""
|
||||
from backend.search import get_all_tags
|
||||
|
||||
if params.vault and params.vault != "all":
|
||||
ctx.require_vault_access(params.vault)
|
||||
return [{"tag": tag, "count": count} for tag, count in get_all_tags(params.vault).items()]
|
||||
return [{"tag": tag, "count": count} for tag, count in _list_tags(params.vault).items()]
|
||||
|
||||
from backend.indexer import get_vault_names
|
||||
|
||||
@@ -161,6 +108,6 @@ def list_tags(ctx: ToolContext, params: ListTagsInput) -> list[dict[str, Any]]:
|
||||
for name in get_vault_names():
|
||||
if not ctx.has_vault_access(name):
|
||||
continue
|
||||
for tag, count in get_all_tags(name).items():
|
||||
for tag, count in _list_tags(name).items():
|
||||
merged[tag] = merged.get(tag, 0) + count
|
||||
return [{"tag": tag, "count": count} for tag, count in sorted(merged.items(), key=lambda x: -x[1])]
|
||||
|
||||
Reference in New Issue
Block a user