feat: 适配 deepagents 最新版特性
This commit is contained in:
parent
b51df7097e
commit
31a7fc52d0
@ -19,7 +19,7 @@ dependencies = [
|
||||
"chardet>=5.0.0",
|
||||
"colorlog>=6.9.0",
|
||||
"dashscope>=1.23.2",
|
||||
"deepagents>=0.2.5",
|
||||
"deepagents>=0.6.7",
|
||||
"docling>=2.68.0",
|
||||
"docx2txt>=0.9",
|
||||
"httpx>=0.27.0",
|
||||
|
||||
@ -1,6 +1,6 @@
|
||||
from deepagents.backends import CompositeBackend, StateBackend
|
||||
|
||||
from .composite import create_agent_composite_backend
|
||||
from .composite import create_agent_composite_backend, create_agent_filesystem_middleware
|
||||
from .knowledge_base_backend import resolve_visible_knowledge_bases_for_context
|
||||
from .sandbox import (
|
||||
SKILLS_PATH,
|
||||
@ -27,6 +27,7 @@ __all__ = [
|
||||
"StateBackend",
|
||||
"SelectedSkillsReadonlyBackend",
|
||||
"create_agent_composite_backend",
|
||||
"create_agent_filesystem_middleware",
|
||||
"ProvisionerSandboxBackend",
|
||||
"ProvisionerSandboxProvider",
|
||||
"SandboxConnection",
|
||||
|
||||
@ -6,65 +6,86 @@ from deepagents.backends.composite import (
|
||||
_route_for_path,
|
||||
_strip_route_from_pattern,
|
||||
)
|
||||
from deepagents.backends.protocol import FileInfo
|
||||
from deepagents.backends.protocol import FileInfo, GlobResult
|
||||
from deepagents.middleware.filesystem import FilesystemMiddleware
|
||||
|
||||
from yuxi.agents.skills.service import normalize_string_list
|
||||
from yuxi.utils.paths import VIRTUAL_PATH_CONVERSATION_HISTORY, VIRTUAL_PATH_LARGE_TOOL_RESULTS, VIRTUAL_PATH_OUTPUTS
|
||||
|
||||
from .sandbox import ProvisionerSandboxBackend
|
||||
from .skills_backend import SelectedSkillsReadonlyBackend
|
||||
|
||||
|
||||
def _coerce_glob_result(result) -> GlobResult:
|
||||
if isinstance(result, GlobResult):
|
||||
return result
|
||||
return GlobResult(matches=result or [])
|
||||
|
||||
|
||||
class CustomCompositeBackend(CompositeBackend):
|
||||
"""修复 glob_info 路由逻辑的 CompositeBackend。
|
||||
"""修复 glob 路由逻辑的 CompositeBackend。
|
||||
|
||||
修复内容:当 path 不匹配任何路由时应该只搜索 default 后端,
|
||||
而不是错误地遍历所有路由后端搜索。
|
||||
"""
|
||||
|
||||
def glob_info(self, pattern: str, path: str = "/") -> list[FileInfo]:
|
||||
def glob(self, pattern: str, path: str = "/") -> GlobResult:
|
||||
backend, backend_path, route_prefix = _route_for_path(
|
||||
default=self.default,
|
||||
sorted_routes=self.sorted_routes,
|
||||
path=path,
|
||||
)
|
||||
if route_prefix is not None:
|
||||
infos = backend.glob_info(pattern, backend_path)
|
||||
return [_remap_file_info_path(fi, route_prefix) for fi in infos]
|
||||
result = _coerce_glob_result(backend.glob(pattern, backend_path))
|
||||
if result.error:
|
||||
return result
|
||||
return GlobResult(matches=[_remap_file_info_path(fi, route_prefix) for fi in (result.matches or [])])
|
||||
|
||||
# 只在 path 为 None 或 "/" 时搜索所有后端,其他只搜索 default
|
||||
if path is None or path == "/":
|
||||
results: list[FileInfo] = []
|
||||
results.extend(self.default.glob_info(pattern, path))
|
||||
default_result = _coerce_glob_result(self.default.glob(pattern, path))
|
||||
if default_result.error:
|
||||
return default_result
|
||||
results.extend(default_result.matches or [])
|
||||
for route_prefix, backend in self.routes.items():
|
||||
route_pattern = _strip_route_from_pattern(pattern, route_prefix)
|
||||
infos = backend.glob_info(route_pattern, "/")
|
||||
results.extend(_remap_file_info_path(fi, route_prefix) for fi in infos)
|
||||
result = _coerce_glob_result(backend.glob(route_pattern, "/"))
|
||||
if result.error:
|
||||
return result
|
||||
results.extend(_remap_file_info_path(fi, route_prefix) for fi in (result.matches or []))
|
||||
results.sort(key=lambda x: x.get("path", ""))
|
||||
return results
|
||||
return GlobResult(matches=results)
|
||||
|
||||
return self.default.glob_info(pattern, path)
|
||||
return _coerce_glob_result(self.default.glob(pattern, path))
|
||||
|
||||
async def aglob_info(self, pattern: str, path: str = "/") -> list[FileInfo]:
|
||||
async def aglob(self, pattern: str, path: str = "/") -> GlobResult:
|
||||
backend, backend_path, route_prefix = _route_for_path(
|
||||
default=self.default,
|
||||
sorted_routes=self.sorted_routes,
|
||||
path=path,
|
||||
)
|
||||
if route_prefix is not None:
|
||||
infos = await backend.aglob_info(pattern, backend_path)
|
||||
return [_remap_file_info_path(fi, route_prefix) for fi in infos]
|
||||
result = _coerce_glob_result(await backend.aglob(pattern, backend_path))
|
||||
if result.error:
|
||||
return result
|
||||
return GlobResult(matches=[_remap_file_info_path(fi, route_prefix) for fi in (result.matches or [])])
|
||||
|
||||
if path is None or path == "/":
|
||||
results: list[FileInfo] = []
|
||||
results.extend(await self.default.aglob_info(pattern, path))
|
||||
default_result = _coerce_glob_result(await self.default.aglob(pattern, path))
|
||||
if default_result.error:
|
||||
return default_result
|
||||
results.extend(default_result.matches or [])
|
||||
for route_prefix, backend in self.routes.items():
|
||||
route_pattern = _strip_route_from_pattern(pattern, route_prefix)
|
||||
infos = await backend.aglob_info(route_pattern, "/")
|
||||
results.extend(_remap_file_info_path(fi, route_prefix) for fi in infos)
|
||||
result = _coerce_glob_result(await backend.aglob(route_pattern, "/"))
|
||||
if result.error:
|
||||
return result
|
||||
results.extend(_remap_file_info_path(fi, route_prefix) for fi in (result.matches or []))
|
||||
results.sort(key=lambda x: x.get("path", ""))
|
||||
return results
|
||||
return GlobResult(matches=results)
|
||||
|
||||
return await self.default.aglob_info(pattern, path)
|
||||
return _coerce_glob_result(await self.default.aglob(pattern, path))
|
||||
|
||||
|
||||
def _get_readable_skills_from_runtime(runtime) -> list[str]:
|
||||
@ -116,4 +137,15 @@ def create_agent_composite_backend(runtime) -> CompositeBackend:
|
||||
routes={
|
||||
"/skills/": SelectedSkillsReadonlyBackend(selected_slugs=readable_skills),
|
||||
},
|
||||
artifacts_root=VIRTUAL_PATH_OUTPUTS,
|
||||
)
|
||||
|
||||
|
||||
def create_agent_filesystem_middleware(tool_token_limit_before_evict: int | None = None) -> FilesystemMiddleware:
|
||||
middleware = FilesystemMiddleware(
|
||||
backend=create_agent_composite_backend,
|
||||
tool_token_limit_before_evict=tool_token_limit_before_evict,
|
||||
)
|
||||
middleware._large_tool_results_prefix = VIRTUAL_PATH_LARGE_TOOL_RESULTS
|
||||
middleware._conversation_history_prefix = VIRTUAL_PATH_CONVERSATION_HISTORY
|
||||
return middleware
|
||||
|
||||
@ -11,7 +11,11 @@ from deepagents.backends.protocol import (
|
||||
FileDownloadResponse,
|
||||
FileInfo,
|
||||
FileUploadResponse,
|
||||
GlobResult,
|
||||
GrepMatch,
|
||||
GrepResult,
|
||||
LsResult,
|
||||
ReadResult,
|
||||
WriteResult,
|
||||
)
|
||||
from deepagents.backends.sandbox import BaseSandbox
|
||||
@ -19,21 +23,106 @@ from deepagents.backends.sandbox import BaseSandbox
|
||||
from yuxi import config as conf
|
||||
from yuxi.agents.skills.service import sync_thread_readable_skills
|
||||
from yuxi.utils.logging_config import logger
|
||||
from yuxi.utils.paths import (
|
||||
OUTPUTS_DIR_NAME,
|
||||
UPLOADS_DIR_NAME,
|
||||
VIRTUAL_PATH_PREFIX,
|
||||
VIRTUAL_SKILLS_PATH,
|
||||
WORKSPACE_DIR_NAME,
|
||||
)
|
||||
|
||||
from .provider import get_sandbox_provider, sandbox_id_for_thread
|
||||
|
||||
|
||||
_USER_DATA_ROOT = "/" + VIRTUAL_PATH_PREFIX.strip("/")
|
||||
_WORKSPACE_ROOT = f"{_USER_DATA_ROOT}/{WORKSPACE_DIR_NAME}"
|
||||
_UPLOADS_ROOT = f"{_USER_DATA_ROOT}/{UPLOADS_DIR_NAME}"
|
||||
_OUTPUTS_ROOT = f"{_USER_DATA_ROOT}/{OUTPUTS_DIR_NAME}"
|
||||
_SKILLS_ROOT = "/" + VIRTUAL_SKILLS_PATH.strip("/")
|
||||
_READABLE_ROOTS = (_USER_DATA_ROOT, _SKILLS_ROOT)
|
||||
_WRITABLE_ROOTS = (_WORKSPACE_ROOT, _OUTPUTS_ROOT)
|
||||
|
||||
|
||||
def _normalize_path(path: str) -> str:
|
||||
raw = str(path or "").strip()
|
||||
if not raw:
|
||||
raise ValueError("path is required")
|
||||
normalized = "/" + raw.lstrip("/")
|
||||
pure = PurePosixPath(normalized)
|
||||
if not raw.startswith("/"):
|
||||
raise ValueError("path must start with /")
|
||||
pure = PurePosixPath(raw)
|
||||
if ".." in pure.parts:
|
||||
raise ValueError("path traversal is not allowed")
|
||||
return str(pure)
|
||||
|
||||
|
||||
def _is_same_or_child(path: str, root: str) -> bool:
|
||||
root = root.rstrip("/") or "/"
|
||||
if root == "/":
|
||||
return path == "/" or path.startswith("/")
|
||||
return path == root or path.startswith(f"{root}/")
|
||||
|
||||
|
||||
def _path_overlaps_root(path: str, root: str) -> bool:
|
||||
return _is_same_or_child(path, root) or _is_same_or_child(root, path)
|
||||
|
||||
|
||||
def _can_read_path(path: str) -> bool:
|
||||
return any(_is_same_or_child(path, root) for root in _READABLE_ROOTS)
|
||||
|
||||
|
||||
def _can_list_path(path: str) -> bool:
|
||||
return any(_path_overlaps_root(path, root) for root in _READABLE_ROOTS)
|
||||
|
||||
|
||||
def _can_write_path(path: str) -> bool:
|
||||
return any(_is_same_or_child(path, root) for root in _WRITABLE_ROOTS)
|
||||
|
||||
|
||||
def _readable_search_paths(path: str) -> list[str]:
|
||||
if _can_read_path(path):
|
||||
return [path]
|
||||
return [root for root in _READABLE_ROOTS if _is_same_or_child(root, path)]
|
||||
|
||||
|
||||
def _glob_for_search_root(pattern: str, root: str) -> str:
|
||||
bare_pattern = str(pattern or "*").lstrip("/")
|
||||
bare_root = root.strip("/")
|
||||
if bare_pattern == bare_root:
|
||||
return "*"
|
||||
root_prefix = f"{bare_root}/"
|
||||
if bare_pattern.startswith(root_prefix):
|
||||
return bare_pattern[len(root_prefix) :] or "*"
|
||||
return pattern
|
||||
|
||||
|
||||
def _filter_readable_infos(infos: list[FileInfo]) -> list[FileInfo]:
|
||||
result: list[FileInfo] = []
|
||||
for info in infos:
|
||||
try:
|
||||
path = _normalize_path(info.get("path", ""))
|
||||
except ValueError:
|
||||
continue
|
||||
if _can_list_path(path):
|
||||
result.append(info)
|
||||
return result
|
||||
|
||||
|
||||
def _filter_readable_matches(matches: list[GrepMatch]) -> list[GrepMatch]:
|
||||
result: list[GrepMatch] = []
|
||||
for match in matches:
|
||||
try:
|
||||
path = _normalize_path(match.get("path", ""))
|
||||
except ValueError:
|
||||
continue
|
||||
if _can_read_path(path):
|
||||
result.append(match)
|
||||
return result
|
||||
|
||||
|
||||
def _permission_error(operation: str, path: str) -> str:
|
||||
return f"permission denied for {operation} on '{path}'"
|
||||
|
||||
|
||||
def _describe_read_error(file_path: str, exc: Exception) -> str:
|
||||
if isinstance(exc, FileNotFoundError):
|
||||
return f"Error: File '{file_path}' not found"
|
||||
@ -111,8 +200,8 @@ class ProvisionerSandboxBackend(BaseSandbox):
|
||||
This helper is the single normalization point used by read(), edit(), and
|
||||
download_files() so all read paths share the same transport semantics.
|
||||
"""
|
||||
start_line = max(0, int(offset)) if offset else None
|
||||
end_line = (start_line + int(limit)) if limit and start_line is not None else None
|
||||
start_line = max(0, int(offset))
|
||||
end_line = start_line + int(limit) if limit is not None else None
|
||||
|
||||
result = self._get_client().file.read_file(
|
||||
file=path,
|
||||
@ -138,35 +227,27 @@ class ProvisionerSandboxBackend(BaseSandbox):
|
||||
file_path: str,
|
||||
offset: int = 0,
|
||||
limit: int = 2000,
|
||||
) -> str:
|
||||
"""Read file content via the sandbox file API and render a text view.
|
||||
|
||||
This stays on top of _read_binary() so the backend has one consistent
|
||||
read path for base64 transport, raw bytes, and text-like responses.
|
||||
"""
|
||||
) -> ReadResult:
|
||||
"""Read allowed file content via the sandbox file API."""
|
||||
try:
|
||||
normalized_path = _normalize_path(file_path)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return _describe_read_error(file_path, exc)
|
||||
start = max(0, int(offset))
|
||||
return ReadResult(error=f"Invalid path '{file_path}': {exc}")
|
||||
if not _can_read_path(normalized_path):
|
||||
return ReadResult(error=_permission_error("read", normalized_path))
|
||||
|
||||
try:
|
||||
content = self._read_binary(normalized_path, offset=offset, limit=limit)
|
||||
if _looks_like_binary(content):
|
||||
content = self._read_binary(normalized_path)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return _describe_read_error(file_path, exc)
|
||||
|
||||
if not content:
|
||||
return "System reminder: File exists but has empty contents"
|
||||
error = _describe_read_error(file_path, exc)
|
||||
return ReadResult(error=error.removeprefix("Error: "))
|
||||
|
||||
if _looks_like_binary(content):
|
||||
return f"Error: File '{file_path}' is binary and cannot be rendered as text"
|
||||
return ReadResult(file_data={"content": base64.b64encode(content).decode("ascii"), "encoding": "base64"})
|
||||
|
||||
text = content.decode("utf-8")
|
||||
if not text:
|
||||
return ""
|
||||
|
||||
lines = text.splitlines()
|
||||
return "\n".join(f"{start + idx + 1:6d}\t{line}" for idx, line in enumerate(lines))
|
||||
return ReadResult(file_data={"content": content.decode("utf-8"), "encoding": "utf-8"})
|
||||
|
||||
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
|
||||
"""Execute a shell command in the sandbox.
|
||||
@ -198,13 +279,19 @@ class ProvisionerSandboxBackend(BaseSandbox):
|
||||
logger.error(f"Sandbox execute failed for thread {self._thread_id}: {exc}")
|
||||
return ExecuteResponse(output=f"Error: {exc}", exit_code=1, truncated=False)
|
||||
|
||||
def ls_info(self, path: str) -> list[FileInfo]:
|
||||
"""List direct children of a sandbox path with lightweight metadata."""
|
||||
normalized_path = _normalize_path(path)
|
||||
def ls(self, path: str) -> LsResult:
|
||||
"""List direct children of an allowed sandbox path with lightweight metadata."""
|
||||
try:
|
||||
normalized_path = _normalize_path(path)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return LsResult(error=f"Invalid path '{path}': {exc}")
|
||||
if not _can_list_path(normalized_path):
|
||||
return LsResult(error=_permission_error("read", normalized_path))
|
||||
|
||||
try:
|
||||
result = self._get_client().file.list_path(path=normalized_path, recursive=False, include_size=True)
|
||||
except Exception: # noqa: BLE001
|
||||
return []
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return LsResult(error=str(exc) or f"Failed to list '{path}'")
|
||||
|
||||
entries = result.data.files or []
|
||||
infos: list[FileInfo] = []
|
||||
@ -225,7 +312,7 @@ class ProvisionerSandboxBackend(BaseSandbox):
|
||||
elif isinstance(modified_time, (int, float)):
|
||||
info["modified_at"] = datetime.fromtimestamp(modified_time).isoformat()
|
||||
infos.append(info)
|
||||
return infos
|
||||
return LsResult(entries=_filter_readable_infos(infos))
|
||||
|
||||
def write(self, file_path: str, content: str) -> WriteResult:
|
||||
"""Write a new text file.
|
||||
@ -233,7 +320,12 @@ class ProvisionerSandboxBackend(BaseSandbox):
|
||||
This method is intentionally text-only. Binary payloads should go through
|
||||
upload_files(), which uses base64 encoding for the sandbox file API.
|
||||
"""
|
||||
normalized_path = _normalize_path(file_path)
|
||||
try:
|
||||
normalized_path = _normalize_path(file_path)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return WriteResult(error=f"Error: Invalid path '{file_path}': {exc}")
|
||||
if not _can_write_path(normalized_path):
|
||||
return WriteResult(error=f"Error: {_permission_error('write', normalized_path)}")
|
||||
if not isinstance(content, str):
|
||||
return WriteResult(error="Error: write() only supports text content; use upload_files() for binary data")
|
||||
try:
|
||||
@ -250,7 +342,7 @@ class ProvisionerSandboxBackend(BaseSandbox):
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return WriteResult(error=str(exc) or f"Failed to write file '{file_path}'")
|
||||
|
||||
return WriteResult(path=normalized_path, files_update=None)
|
||||
return WriteResult(path=normalized_path)
|
||||
|
||||
def edit(
|
||||
self,
|
||||
@ -264,7 +356,12 @@ class ProvisionerSandboxBackend(BaseSandbox):
|
||||
This method operates on UTF-8-decoded text content only. Binary files
|
||||
are not supported here and should be handled via download/upload flows.
|
||||
"""
|
||||
normalized_path = _normalize_path(file_path)
|
||||
try:
|
||||
normalized_path = _normalize_path(file_path)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return EditResult(error=f"Error: Invalid path '{file_path}': {exc}")
|
||||
if not _can_write_path(normalized_path):
|
||||
return EditResult(error=f"Error: {_permission_error('write', normalized_path)}")
|
||||
|
||||
# Check if old_string exists
|
||||
try:
|
||||
@ -298,44 +395,59 @@ class ProvisionerSandboxBackend(BaseSandbox):
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return EditResult(error=f"Error editing file: {exc}")
|
||||
|
||||
return EditResult(path=normalized_path, files_update=None, occurrences=count if replace_all else 1)
|
||||
return EditResult(path=normalized_path, occurrences=count if replace_all else 1)
|
||||
|
||||
def grep_raw(
|
||||
def grep(
|
||||
self,
|
||||
pattern: str,
|
||||
path: str | None = None,
|
||||
glob: str | None = None,
|
||||
) -> list[GrepMatch] | str:
|
||||
"""Search file contents under a path and return raw line matches.
|
||||
|
||||
The sandbox file API is used directly with fixed-string matching and an
|
||||
optional include glob.
|
||||
"""
|
||||
search_path = _normalize_path(path or "/")
|
||||
|
||||
) -> GrepResult:
|
||||
"""Search allowed sandbox paths for literal text."""
|
||||
try:
|
||||
return super().grep_raw(pattern=pattern, path=search_path, glob=glob)
|
||||
|
||||
normalized_path = _normalize_path(path or "/")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return str(exc)
|
||||
return GrepResult(error=f"Invalid path '{path or '/'}': {exc}")
|
||||
|
||||
def glob_info(self, pattern: str, path: str = "/") -> list[FileInfo]:
|
||||
"""Return files matching a glob pattern with optional metadata."""
|
||||
normalized_path = _normalize_path(path)
|
||||
search_paths = _readable_search_paths(normalized_path)
|
||||
if not search_paths:
|
||||
return GrepResult(error=_permission_error("read", normalized_path))
|
||||
|
||||
matches: list[GrepMatch] = []
|
||||
for search_path in search_paths:
|
||||
result = super().grep(pattern=pattern, path=search_path, glob=glob)
|
||||
if result.error:
|
||||
return result
|
||||
matches.extend(result.matches or [])
|
||||
return GrepResult(matches=_filter_readable_matches(matches))
|
||||
|
||||
def glob(self, pattern: str, path: str = "/") -> GlobResult:
|
||||
"""Return files matching a glob pattern under allowed sandbox paths."""
|
||||
try:
|
||||
# return super().glob_info(pattern=pattern, path=path)
|
||||
result = self._get_client().file.find_files(
|
||||
path=normalized_path,
|
||||
glob=pattern,
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
return []
|
||||
normalized_path = _normalize_path(path)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return GlobResult(error=f"Invalid path '{path}': {exc}")
|
||||
if ".." in PurePosixPath(str(pattern or "")).parts:
|
||||
return GlobResult(error="Invalid glob pattern: path traversal is not allowed")
|
||||
|
||||
search_paths = _readable_search_paths(normalized_path)
|
||||
if not search_paths:
|
||||
return GlobResult(error=_permission_error("read", normalized_path))
|
||||
|
||||
infos: list[FileInfo] = []
|
||||
for file_path in result.data.files or []:
|
||||
infos.append({"path": file_path})
|
||||
return infos
|
||||
for search_path in search_paths:
|
||||
try:
|
||||
result = self._get_client().file.find_files(
|
||||
path=search_path,
|
||||
glob=_glob_for_search_root(pattern, search_path),
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
return GlobResult(error=str(exc) or f"Failed to glob '{path}'")
|
||||
for file_path in result.data.files or []:
|
||||
infos.append({"path": file_path})
|
||||
infos = _filter_readable_infos(infos)
|
||||
infos.sort(key=lambda item: item.get("path", ""))
|
||||
return GlobResult(matches=infos)
|
||||
|
||||
def upload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]:
|
||||
"""Upload binary or text file payloads via the sandbox file API.
|
||||
@ -347,6 +459,9 @@ class ProvisionerSandboxBackend(BaseSandbox):
|
||||
for path, content in files:
|
||||
try:
|
||||
normalized_path = _normalize_path(path)
|
||||
if not _can_write_path(normalized_path):
|
||||
responses.append(FileUploadResponse(path=normalized_path, error="permission_denied"))
|
||||
continue
|
||||
result = self._get_client().file.write_file(
|
||||
file=normalized_path,
|
||||
content=base64.b64encode(content).decode("ascii"),
|
||||
@ -380,6 +495,11 @@ class ProvisionerSandboxBackend(BaseSandbox):
|
||||
for path in paths:
|
||||
try:
|
||||
normalized_path = _normalize_path(path)
|
||||
if not _can_read_path(normalized_path):
|
||||
responses.append(
|
||||
FileDownloadResponse(path=normalized_path, content=None, error="permission_denied")
|
||||
)
|
||||
continue
|
||||
content = self._read_binary(normalized_path)
|
||||
responses.append(FileDownloadResponse(path=normalized_path, content=content, error=None))
|
||||
except PermissionError:
|
||||
|
||||
@ -1,10 +1,20 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import PurePosixPath
|
||||
from typing import Any
|
||||
|
||||
from deepagents.backends import FilesystemBackend
|
||||
from deepagents.backends.protocol import EditResult, FileDownloadResponse, FileInfo, FileUploadResponse, WriteResult
|
||||
from deepagents.backends.protocol import (
|
||||
EditResult,
|
||||
FileDownloadResponse,
|
||||
FileInfo,
|
||||
FileUploadResponse,
|
||||
GlobResult,
|
||||
GrepMatch,
|
||||
GrepResult,
|
||||
LsResult,
|
||||
ReadResult,
|
||||
WriteResult,
|
||||
)
|
||||
|
||||
from yuxi.agents.skills.service import get_skills_root_dir, is_valid_skill_slug
|
||||
|
||||
@ -40,53 +50,62 @@ class SelectedSkillsReadonlyBackend(FilesystemBackend):
|
||||
slug = self._extract_slug(file_path)
|
||||
return slug is not None and slug in self._selected_slugs
|
||||
|
||||
def ls_info(self, path: str) -> list[FileInfo]:
|
||||
def _filter_infos(self, infos: list[FileInfo]) -> list[FileInfo]:
|
||||
return [item for item in infos if self._extract_slug(item.get("path", "")) in self._selected_slugs]
|
||||
|
||||
def _filter_matches(self, matches: list[GrepMatch]) -> list[GrepMatch]:
|
||||
return [item for item in matches if self._extract_slug(item.get("path", "")) in self._selected_slugs]
|
||||
|
||||
def ls(self, path: str) -> LsResult:
|
||||
if not self._selected_slugs:
|
||||
return []
|
||||
return LsResult(entries=[])
|
||||
|
||||
normalized = (path or "/").strip() or "/"
|
||||
if not self._is_allowed_path(normalized):
|
||||
return []
|
||||
return LsResult(error="Access denied: path is outside selected skills.")
|
||||
|
||||
infos = super().ls_info(normalized)
|
||||
if normalized == "/":
|
||||
result = []
|
||||
for item in infos:
|
||||
slug = self._extract_slug(item.get("path", ""))
|
||||
if slug in self._selected_slugs:
|
||||
result.append(item)
|
||||
result = super().ls(normalized)
|
||||
if result.error:
|
||||
return result
|
||||
return infos
|
||||
infos = result.entries or []
|
||||
if normalized == "/":
|
||||
infos = self._filter_infos(infos)
|
||||
return LsResult(entries=infos)
|
||||
|
||||
def read(self, file_path: str, offset: int = 0, limit: int = 2000) -> str:
|
||||
def read(self, file_path: str, offset: int = 0, limit: int = 2000) -> ReadResult:
|
||||
if not self._is_allowed_file(file_path):
|
||||
return "Access denied: file is outside selected skills."
|
||||
return ReadResult(error="Access denied: file is outside selected skills.")
|
||||
return super().read(file_path, offset=offset, limit=limit)
|
||||
|
||||
def grep_raw(self, pattern: str, path: str | None = None, glob: str | None = None) -> list[dict] | str:
|
||||
def grep(self, pattern: str, path: str | None = None, glob: str | None = None) -> GrepResult:
|
||||
if not self._selected_slugs:
|
||||
return []
|
||||
return GrepResult(matches=[])
|
||||
|
||||
if path is not None:
|
||||
if not self._is_allowed_path(path):
|
||||
return "Access denied: path is outside selected skills."
|
||||
return super().grep_raw(pattern=pattern, path=path, glob=glob)
|
||||
return GrepResult(error="Access denied: path is outside selected skills.")
|
||||
result = super().grep(pattern=pattern, path=path, glob=glob)
|
||||
if result.error:
|
||||
return result
|
||||
return GrepResult(matches=self._filter_matches(result.matches or []))
|
||||
|
||||
matches: list[dict[str, Any]] = []
|
||||
matches: list[GrepMatch] = []
|
||||
for slug in sorted(self._selected_slugs):
|
||||
result = super().grep_raw(pattern=pattern, path=f"/{slug}", glob=glob)
|
||||
if isinstance(result, str):
|
||||
result = super().grep(pattern=pattern, path=f"/{slug}", glob=glob)
|
||||
if result.error:
|
||||
continue
|
||||
matches.extend(result)
|
||||
return matches
|
||||
matches.extend(result.matches or [])
|
||||
return GrepResult(matches=matches)
|
||||
|
||||
def glob_info(self, pattern: str, path: str = "/") -> list[FileInfo]:
|
||||
def glob(self, pattern: str, path: str = "/") -> GlobResult:
|
||||
if not self._selected_slugs:
|
||||
return []
|
||||
return GlobResult(matches=[])
|
||||
if not self._is_allowed_path(path):
|
||||
return []
|
||||
infos = super().glob_info(pattern=pattern, path=path)
|
||||
return [item for item in infos if self._extract_slug(item.get("path", "")) in self._selected_slugs]
|
||||
return GlobResult(error="Access denied: path is outside selected skills.")
|
||||
result = super().glob(pattern=pattern, path=path)
|
||||
if result.error:
|
||||
return result
|
||||
return GlobResult(matches=self._filter_infos(result.matches or []))
|
||||
|
||||
def write(self, file_path: str, content: str) -> WriteResult:
|
||||
return WriteResult(error="Skills path is read-only.")
|
||||
|
||||
@ -1,19 +1,18 @@
|
||||
from deepagents.middleware.filesystem import FilesystemMiddleware
|
||||
from deepagents.middleware.patch_tool_calls import PatchToolCallsMiddleware
|
||||
from deepagents.middleware.subagents import SubAgentMiddleware
|
||||
from langchain.agents import create_agent
|
||||
from langchain.agents.middleware import ModelRetryMiddleware, TodoListMiddleware
|
||||
|
||||
from yuxi.agents import BaseAgent, BaseState, load_chat_model
|
||||
from yuxi.agents.backends import create_agent_composite_backend
|
||||
from yuxi.agents.backends import create_agent_composite_backend, create_agent_filesystem_middleware
|
||||
from yuxi.agents.context import prepare_agent_runtime_context
|
||||
from yuxi.agents.middlewares import (
|
||||
SummaryOffloadMiddleware,
|
||||
create_summary_middleware,
|
||||
save_attachments_to_fs,
|
||||
)
|
||||
from yuxi.agents.middlewares.knowledge_base_middleware import KnowledgeBaseMiddleware
|
||||
from yuxi.agents.middlewares.skills_middleware import SkillsMiddleware
|
||||
from yuxi.agents.subagents.service import get_subagents_from_slugs
|
||||
from yuxi.agents.subagents.service import build_subagent_middleware_specs, get_subagents_from_slugs
|
||||
from yuxi.agents.toolkits.service import resolve_configured_runtime_tools
|
||||
|
||||
from .prompt import TODO_MID_PROMPT, build_prompt_with_context
|
||||
@ -22,30 +21,35 @@ from .prompt import TODO_MID_PROMPT, build_prompt_with_context
|
||||
async def _build_middlewares(context):
|
||||
"""构建中间件列表"""
|
||||
# summary middleware
|
||||
# 主 Agent 上下文优化:90k tokens 触发压缩(128k context window 的 70%)
|
||||
summary_middleware = SummaryOffloadMiddleware(
|
||||
# 主 Agent 上下文优化:默认 100k tokens 触发压缩,保留最近 50%
|
||||
summary_trigger_tokens = getattr(context, "summary_threshold", 100) * 1024
|
||||
summary_middleware = create_summary_middleware(
|
||||
model=load_chat_model(fully_specified_name=context.model),
|
||||
trigger=("tokens", getattr(context, "summary_threshold", 100) * 1024),
|
||||
trigger=("tokens", summary_trigger_tokens),
|
||||
keep=("tokens", summary_trigger_tokens // 2),
|
||||
trim_tokens_to_summarize=4000,
|
||||
summary_offload_threshold=500,
|
||||
max_retention_ratio=0.5,
|
||||
)
|
||||
|
||||
# subagents
|
||||
subagents = await get_subagents_from_slugs(context.subagents)
|
||||
default_subagent_middleware = [
|
||||
create_agent_filesystem_middleware(tool_token_limit_before_evict=500), # 文件系统后端
|
||||
PatchToolCallsMiddleware(),
|
||||
summary_middleware,
|
||||
]
|
||||
subagents_middleware = SubAgentMiddleware(
|
||||
default_model=load_chat_model(fully_specified_name=context.subagents_model),
|
||||
subagents=subagents,
|
||||
general_purpose_agent=True,
|
||||
default_middleware=[
|
||||
FilesystemMiddleware(backend=create_agent_composite_backend), # 文件系统后端
|
||||
PatchToolCallsMiddleware(),
|
||||
summary_middleware,
|
||||
],
|
||||
backend=create_agent_composite_backend,
|
||||
subagents=build_subagent_middleware_specs(
|
||||
subagents,
|
||||
default_model=load_chat_model(fully_specified_name=context.subagents_model),
|
||||
default_middleware=default_subagent_middleware,
|
||||
model_loader=load_chat_model,
|
||||
),
|
||||
state_schema=BaseState,
|
||||
)
|
||||
# all middlewares
|
||||
middlewares = [
|
||||
FilesystemMiddleware(backend=create_agent_composite_backend), # 文件系统后端
|
||||
create_agent_filesystem_middleware(tool_token_limit_before_evict=500), # 文件系统后端
|
||||
save_attachments_to_fs, # 附件注入提示词
|
||||
KnowledgeBaseMiddleware(), # 知识库工具
|
||||
SkillsMiddleware(), # Skills 中间件(提示词注入、依赖展开、动态激活)
|
||||
@ -63,6 +67,7 @@ class ChatbotAgent(BaseAgent):
|
||||
name = "智能助手"
|
||||
description = "基础的对话机器人,可以回答问题,可在配置中启用需要的工具。"
|
||||
capabilities = ["file_upload", "files"] # 支持文件上传功能
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
||||
@ -1,6 +1,5 @@
|
||||
import os
|
||||
|
||||
from deepagents.middleware.filesystem import FilesystemMiddleware
|
||||
from deepagents.middleware.patch_tool_calls import PatchToolCallsMiddleware
|
||||
from deepagents.middleware.subagents import SubAgentMiddleware
|
||||
from langchain.agents import create_agent
|
||||
@ -10,16 +9,16 @@ from langchain.agents.middleware import (
|
||||
)
|
||||
|
||||
from yuxi.agents import BaseAgent, BaseState, load_chat_model
|
||||
from yuxi.agents.backends import create_agent_composite_backend
|
||||
from yuxi.agents.backends import create_agent_composite_backend, create_agent_filesystem_middleware
|
||||
from yuxi.agents.context import prepare_agent_runtime_context
|
||||
from yuxi.agents.middlewares import (
|
||||
SummaryOffloadMiddleware,
|
||||
create_summary_middleware,
|
||||
save_attachments_to_fs,
|
||||
)
|
||||
from yuxi.agents.middlewares.knowledge_base_middleware import KnowledgeBaseMiddleware
|
||||
from yuxi.agents.middlewares.skills_middleware import SkillsMiddleware
|
||||
from yuxi.agents.toolkits.buildin.tools import _create_tavily_search
|
||||
from yuxi.agents.subagents.service import get_subagents_from_slugs
|
||||
from yuxi.agents.subagents.service import build_subagent_middleware_specs, get_subagents_from_slugs
|
||||
from yuxi.agents.toolkits.service import resolve_configured_runtime_tools
|
||||
from yuxi.utils import logger
|
||||
from yuxi.utils.datetime_utils import shanghai_now
|
||||
@ -65,31 +64,35 @@ class DeepAgent(BaseAgent):
|
||||
# 从数据库加载 subagent specs(工具名称已解析)
|
||||
user_subagents = await get_subagents_from_slugs(context.subagents)
|
||||
|
||||
# 主 Agent 上下文优化:90k tokens 触发压缩(128k context window 的 70%)
|
||||
summary_middleware = SummaryOffloadMiddleware(
|
||||
# 主 Agent 上下文优化:90k tokens 触发压缩,保留最近 50%
|
||||
summary_middleware = create_summary_middleware(
|
||||
model=model,
|
||||
trigger=("tokens", 90000),
|
||||
keep=("tokens", 45000),
|
||||
trim_tokens_to_summarize=4000,
|
||||
summary_offload_threshold=500,
|
||||
max_retention_ratio=0.5,
|
||||
)
|
||||
|
||||
default_subagent_middleware = [
|
||||
create_agent_filesystem_middleware(tool_token_limit_before_evict=500), # 文件系统后端
|
||||
PatchToolCallsMiddleware(),
|
||||
summary_middleware,
|
||||
# 子 Agent 搜索工具限制:tavily_search 最多 8 次
|
||||
ToolCallLimitMiddleware(
|
||||
tool_name="tavily_search",
|
||||
run_limit=8,
|
||||
exit_behavior="continue",
|
||||
),
|
||||
]
|
||||
subagents_middleware = SubAgentMiddleware(
|
||||
default_model=sub_model,
|
||||
default_tools=search_tools,
|
||||
subagents=user_subagents,
|
||||
default_middleware=[
|
||||
FilesystemMiddleware(backend=create_agent_composite_backend), # 文件系统后端
|
||||
PatchToolCallsMiddleware(),
|
||||
summary_middleware,
|
||||
# 子 Agent 搜索工具限制:tavily_search 最多 8 次
|
||||
ToolCallLimitMiddleware(
|
||||
tool_name="tavily_search",
|
||||
run_limit=8,
|
||||
exit_behavior="continue",
|
||||
),
|
||||
],
|
||||
general_purpose_agent=True,
|
||||
backend=create_agent_composite_backend,
|
||||
subagents=build_subagent_middleware_specs(
|
||||
user_subagents,
|
||||
default_model=sub_model,
|
||||
default_tools=search_tools,
|
||||
default_middleware=default_subagent_middleware,
|
||||
model_loader=load_chat_model,
|
||||
),
|
||||
state_schema=BaseState,
|
||||
)
|
||||
|
||||
# 使用 create_deep_agent 创建深度智能体
|
||||
@ -98,7 +101,7 @@ class DeepAgent(BaseAgent):
|
||||
tools=await resolve_configured_runtime_tools(context),
|
||||
system_prompt=system_prompt,
|
||||
middleware=[
|
||||
FilesystemMiddleware(backend=create_agent_composite_backend), # 文件系统后端
|
||||
create_agent_filesystem_middleware(tool_token_limit_before_evict=500), # 文件系统后端
|
||||
SkillsMiddleware(), # Skills 中间件(提示词注入、依赖展开、动态激活)
|
||||
save_attachments_to_fs, # 附件注入提示词
|
||||
TodoListMiddleware(system_prompt="任务结束前,应该检查维护的待办事项列表是否结束。"),
|
||||
|
||||
@ -1,14 +1,13 @@
|
||||
from .attachment_middleware import inject_attachment_context, save_attachments_to_fs
|
||||
from .context_middlewares import context_aware_prompt, context_based_model
|
||||
from .dynamic_tool_middleware import DynamicToolMiddleware
|
||||
from .summary_middleware import SummaryOffloadMiddleware, create_summary_offload_middleware
|
||||
from .summary_middleware import create_summary_middleware
|
||||
|
||||
__all__ = [
|
||||
"DynamicToolMiddleware",
|
||||
"SummaryOffloadMiddleware",
|
||||
"context_aware_prompt",
|
||||
"context_based_model",
|
||||
"create_summary_offload_middleware",
|
||||
"create_summary_middleware",
|
||||
"inject_attachment_context", # 已废弃,使用 save_attachments_to_fs
|
||||
"save_attachments_to_fs",
|
||||
]
|
||||
|
||||
@ -6,6 +6,7 @@ from collections.abc import Callable
|
||||
from pathlib import PurePosixPath
|
||||
from typing import Annotated, Any, NotRequired, TypedDict
|
||||
|
||||
from deepagents.middleware._utils import append_to_system_message
|
||||
from deepagents.middleware.skills import SKILLS_SYSTEM_PROMPT
|
||||
from langchain.agents import AgentState
|
||||
from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse
|
||||
@ -197,42 +198,22 @@ class SkillsMiddleware(AgentMiddleware):
|
||||
self.enable_skills_prompt = enable_skills_prompt
|
||||
self.skills_sources_for_prompt = skills_sources_for_prompt or ["/home/gem/skills/"]
|
||||
|
||||
async def abefore_agent(self, state: SkillsState, runtime) -> dict[str, Any] | None:
|
||||
"""在 agent 执行前注入 skills 提示词"""
|
||||
runtime_context = runtime.context
|
||||
|
||||
# 检查是否需要注入
|
||||
if not self.enable_skills_prompt:
|
||||
return None
|
||||
if getattr(runtime_context, "_skills_prompt_injected", False):
|
||||
return None
|
||||
|
||||
prompt_skills = getattr(runtime_context, "_prompt_skills", None)
|
||||
if not isinstance(prompt_skills, list):
|
||||
return None
|
||||
|
||||
prompt_skills = normalize_string_list(prompt_skills)
|
||||
if not prompt_skills:
|
||||
return None
|
||||
|
||||
# 收集提示词元数据并构建提示段
|
||||
skills_meta = self._collect_prompt_metadata(prompt_skills, runtime_context)
|
||||
skills_section = self._build_skills_section(skills_meta)
|
||||
|
||||
# 注入提示词
|
||||
base_prompt = getattr(runtime_context, "system_prompt", "") or ""
|
||||
merged_prompt = f"{base_prompt}\n\n{skills_section}" if base_prompt else skills_section
|
||||
setattr(runtime_context, "system_prompt", merged_prompt)
|
||||
setattr(runtime_context, "_skills_prompt_injected", True)
|
||||
|
||||
return None
|
||||
|
||||
async def awrap_model_call(
|
||||
self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse]
|
||||
) -> ModelResponse:
|
||||
"""包装模型调用,处理动态激活和依赖展开"""
|
||||
"""包装模型调用,处理 skills 提示词注入、动态激活和依赖展开"""
|
||||
runtime_context = request.runtime.context
|
||||
|
||||
if self.enable_skills_prompt:
|
||||
prompt_skills = getattr(runtime_context, "_prompt_skills", None)
|
||||
if isinstance(prompt_skills, list):
|
||||
prompt_skills = normalize_string_list(prompt_skills)
|
||||
if prompt_skills:
|
||||
skills_meta = self._collect_prompt_metadata(prompt_skills, runtime_context)
|
||||
skills_section = self._build_skills_section(skills_meta)
|
||||
system_message = append_to_system_message(getattr(request, "system_message", None), skills_section)
|
||||
request = request.override(system_message=system_message)
|
||||
|
||||
state = request.state if isinstance(request.state, dict) else {}
|
||||
activated = state.get("activated_skills", []) or []
|
||||
if not isinstance(activated, list):
|
||||
@ -473,5 +454,6 @@ class SkillsMiddleware(AgentMiddleware):
|
||||
skills_list = self._format_skills_list(skills_meta)
|
||||
return SKILLS_SYSTEM_PROMPT.format(
|
||||
skills_locations=skills_locations,
|
||||
skills_load_warnings="",
|
||||
skills_list=skills_list,
|
||||
)
|
||||
|
||||
@ -1,817 +1,30 @@
|
||||
"""Summary + ToolResult Offload Middleware.
|
||||
|
||||
基于 LangChain SummarizationMiddleware 实现,额外添加工具结果卸载功能:
|
||||
- 保留原有的 summary 历史记录逻辑
|
||||
- 将 ToolMessage 的 results 卸载到虚拟文件系统(默认 > 1k 字符)
|
||||
"""
|
||||
"""Yuxi adapter for DeepAgents conversation summarization middleware."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterable, Mapping
|
||||
from functools import partial
|
||||
from typing import Any, Literal, cast, override
|
||||
from deepagents.middleware.summarization import SummarizationMiddleware
|
||||
from langchain.agents.middleware.summarization import ContextSize
|
||||
from langchain.chat_models import BaseChatModel
|
||||
|
||||
from langchain.agents import AgentState
|
||||
from langchain.agents.middleware import AgentMiddleware
|
||||
from langchain.chat_models import BaseChatModel, init_chat_model
|
||||
from langchain_core.messages import (
|
||||
AIMessage,
|
||||
AnyMessage,
|
||||
HumanMessage,
|
||||
MessageLikeRepresentation,
|
||||
RemoveMessage,
|
||||
ToolMessage,
|
||||
)
|
||||
from langchain_core.messages.utils import (
|
||||
count_tokens_approximately,
|
||||
get_buffer_string,
|
||||
trim_messages,
|
||||
)
|
||||
from langgraph.graph.message import REMOVE_ALL_MESSAGES
|
||||
from langgraph.runtime import Runtime
|
||||
from yuxi.agents.backends.composite import create_agent_composite_backend
|
||||
from yuxi.utils.paths import VIRTUAL_PATH_CONVERSATION_HISTORY, VIRTUAL_PATH_LARGE_TOOL_RESULTS
|
||||
|
||||
from yuxi.utils.paths import VIRTUAL_PATH_OUTPUTS
|
||||
|
||||
TokenCounter = Callable[[Iterable[MessageLikeRepresentation]], int]
|
||||
|
||||
DEFAULT_SUMMARY_PROMPT = """<role>
|
||||
Context Extraction Assistant
|
||||
</role>
|
||||
|
||||
<primary_objective>
|
||||
Your sole objective in this task is to extract the highest quality/most relevant
|
||||
context from the conversation history below.
|
||||
</primary_objective>
|
||||
|
||||
<objective_information>
|
||||
You're nearing the total number of input tokens you can accept, so you must
|
||||
extract the highest quality/most relevant pieces of information from your conversation
|
||||
history. This context will then overwrite the conversation history presented below.
|
||||
Because of this, ensure the context you extract is only the most important information
|
||||
to your overall goal.
|
||||
</objective_information>
|
||||
|
||||
<instructions>
|
||||
The conversation history below will be replaced with the context you extract in
|
||||
this step. Because of this, you must do your very best to extract and record all
|
||||
of the most important context from the conversation history. You want to ensure
|
||||
that you don't repeat any actions you've already completed, so the context you
|
||||
extract from the conversation history should be focused on the most important
|
||||
information to your overall goal.
|
||||
</instructions>
|
||||
|
||||
The user will message you with the full message history you'll be extracting context
|
||||
from, to then replace. Carefully read over it all, and think deeply about what
|
||||
information is most important to your overall goal that should be saved.
|
||||
|
||||
With all of this in mind, please carefully read over the entire conversation history,
|
||||
and extract the most important and relevant context to replace it so that you can
|
||||
free up space in the conversation history. Respond ONLY with the extracted context.
|
||||
Do not include any additional information, or text before or after the extracted context.
|
||||
|
||||
<messages>
|
||||
Messages to summarize:
|
||||
{messages}
|
||||
</messages>"""
|
||||
|
||||
_DEFAULT_MESSAGES_TO_KEEP = 20
|
||||
_DEFAULT_FALLBACK_MESSAGE_COUNT = 15
|
||||
_OFFLOAD_DIR = "summary_offload"
|
||||
|
||||
ContextFraction = tuple[Literal["fraction"], float]
|
||||
ContextTokens = tuple[Literal["tokens"], int]
|
||||
ContextMessages = tuple[Literal["messages"], int]
|
||||
|
||||
ContextSize = ContextFraction | ContextTokens | ContextMessages
|
||||
|
||||
|
||||
def _get_approximate_token_counter(model: BaseChatModel) -> TokenCounter:
|
||||
"""Tune parameters of approximate token counter based on model type."""
|
||||
if model._llm_type == "anthropic-chat": # noqa: SLF001
|
||||
return partial(count_tokens_approximately, chars_per_token=3.3)
|
||||
return count_tokens_approximately
|
||||
|
||||
|
||||
def _get_content_str(content: Any) -> str | None:
|
||||
"""Convert plain-text ToolMessage content to string for size checking."""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
if len(content) == 1 and isinstance(content[0], dict) and content[0].get("type") == "text":
|
||||
return str(content[0].get("text", ""))
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _format_offload_placeholder(file_path: str, content_sample: str) -> str:
|
||||
"""Format the placeholder message for offloaded content."""
|
||||
return (
|
||||
f"[ToolResultOffloaded]\n\n"
|
||||
f"文件路径: {file_path}\n"
|
||||
f"可以使用 read_file 工具读取完整内容\n\n"
|
||||
f"--- 内容预览 ---\n{content_sample}"
|
||||
)
|
||||
|
||||
|
||||
def _build_offload_file_path(msg: ToolMessage) -> str:
|
||||
"""Build a read_file-compatible virtual path for offloaded tool output."""
|
||||
tool_name = msg.name or "unknown"
|
||||
message_id = msg.id or str(uuid.uuid4())[:8]
|
||||
safe_name = "".join(c if c.isalnum() or c in "-_" else "_" for c in tool_name)
|
||||
return f"{VIRTUAL_PATH_OUTPUTS}/{_OFFLOAD_DIR}/{safe_name}-{message_id}.txt"
|
||||
|
||||
|
||||
def _write_offloaded_content(runtime: Runtime, file_path: str, content: str) -> tuple[bool, dict[str, Any]]:
|
||||
"""Persist offloaded tool output into the active filesystem backend."""
|
||||
from yuxi.agents.backends.composite import create_agent_composite_backend
|
||||
|
||||
backend = create_agent_composite_backend(runtime)
|
||||
result = backend.write(file_path, content)
|
||||
if result.error:
|
||||
return False, {}
|
||||
return True, result.files_update or {}
|
||||
|
||||
|
||||
async def _awrite_offloaded_content(runtime: Runtime, file_path: str, content: str) -> tuple[bool, dict[str, Any]]:
|
||||
"""Async variant of _write_offloaded_content."""
|
||||
from yuxi.agents.backends.composite import create_agent_composite_backend
|
||||
|
||||
backend = create_agent_composite_backend(runtime)
|
||||
result = await backend.awrite(file_path, content)
|
||||
if result.error:
|
||||
return False, {}
|
||||
return True, result.files_update or {}
|
||||
|
||||
|
||||
def _offload_tool_result(
|
||||
msg: ToolMessage,
|
||||
threshold: int,
|
||||
token_counter: TokenCounter,
|
||||
runtime: Runtime,
|
||||
) -> dict[str, Any] | None:
|
||||
"""卸载单个超阈值的工具结果.
|
||||
|
||||
Args:
|
||||
msg: ToolMessage
|
||||
threshold: token 数阈值
|
||||
token_counter: token 计数函数
|
||||
|
||||
Returns:
|
||||
包含 file 更新的字典,如果没有卸载则返回 None
|
||||
"""
|
||||
content = msg.content
|
||||
content_str = _get_content_str(content)
|
||||
|
||||
if content_str is None:
|
||||
return None
|
||||
|
||||
# 计算 token 数
|
||||
msg_tokens = token_counter([msg])
|
||||
if msg_tokens <= threshold:
|
||||
return None
|
||||
|
||||
# 获取工具名称和参数
|
||||
tool_name = msg.name or "unknown"
|
||||
tool_call_id = msg.tool_call_id or ""
|
||||
|
||||
file_path = _build_offload_file_path(msg)
|
||||
|
||||
# 构建文件头部信息
|
||||
header_lines = [
|
||||
"=== Tool Invocation ===",
|
||||
f"Tool: {tool_name}",
|
||||
f"Tool Call ID: {tool_call_id}",
|
||||
"=" * 40,
|
||||
"",
|
||||
]
|
||||
header = "\n".join(header_lines)
|
||||
|
||||
written, files_update = _write_offloaded_content(runtime, file_path, header + content_str)
|
||||
if not written:
|
||||
return None
|
||||
|
||||
# 创建预览内容
|
||||
preview_lines = content_str.splitlines()[:10]
|
||||
content_sample = "\n".join(line[:500] for line in preview_lines)
|
||||
|
||||
# 替换消息内容为占位符
|
||||
msg.content = _format_offload_placeholder(file_path, content_sample)
|
||||
|
||||
return files_update or {}
|
||||
|
||||
|
||||
def _offload_tool_results(
|
||||
messages: list[AnyMessage], threshold: int, token_counter: TokenCounter, runtime: Runtime
|
||||
) -> tuple[dict[str, Any], list[AnyMessage]]:
|
||||
"""扫描消息列表,卸载所有超阈值的工具结果.
|
||||
|
||||
Args:
|
||||
messages: 消息列表
|
||||
threshold: token 数阈值
|
||||
token_counter: token 计数函数
|
||||
|
||||
Returns:
|
||||
tuple[files 更新字典, 被修改的消息列表]
|
||||
"""
|
||||
files_update: dict[str, Any] = {}
|
||||
modified_messages: list[AnyMessage] = []
|
||||
|
||||
for msg in messages:
|
||||
if not isinstance(msg, ToolMessage):
|
||||
continue
|
||||
|
||||
result = _offload_tool_result(msg, threshold, token_counter, runtime)
|
||||
if result is not None:
|
||||
files_update.update(result)
|
||||
modified_messages.append(msg)
|
||||
|
||||
return files_update, modified_messages
|
||||
|
||||
|
||||
async def _aoffload_tool_results(
|
||||
messages: list[AnyMessage], threshold: int, token_counter: TokenCounter, runtime: Runtime
|
||||
) -> tuple[dict[str, Any], list[AnyMessage]]:
|
||||
"""Async variant of _offload_tool_results."""
|
||||
files_update: dict[str, Any] = {}
|
||||
modified_messages: list[AnyMessage] = []
|
||||
|
||||
for msg in messages:
|
||||
if not isinstance(msg, ToolMessage):
|
||||
continue
|
||||
|
||||
content_str = _get_content_str(msg.content)
|
||||
if content_str is None:
|
||||
continue
|
||||
|
||||
msg_tokens = token_counter([msg])
|
||||
if msg_tokens <= threshold:
|
||||
continue
|
||||
|
||||
tool_name = msg.name or "unknown"
|
||||
tool_call_id = msg.tool_call_id or ""
|
||||
file_path = _build_offload_file_path(msg)
|
||||
header_lines = [
|
||||
"=== Tool Invocation ===",
|
||||
f"Tool: {tool_name}",
|
||||
f"Tool Call ID: {tool_call_id}",
|
||||
"=" * 40,
|
||||
"",
|
||||
]
|
||||
header = "\n".join(header_lines)
|
||||
written, result = await _awrite_offloaded_content(runtime, file_path, header + content_str)
|
||||
if not written:
|
||||
continue
|
||||
|
||||
preview_lines = content_str.splitlines()[:10]
|
||||
content_sample = "\n".join(line[:500] for line in preview_lines)
|
||||
msg.content = _format_offload_placeholder(file_path, content_sample)
|
||||
files_update.update(result)
|
||||
modified_messages.append(msg)
|
||||
|
||||
return files_update, modified_messages
|
||||
|
||||
|
||||
class SummaryOffloadMiddleware(AgentMiddleware):
|
||||
"""总结+工具结果卸载中间件.
|
||||
|
||||
基于 LangChain SummarizationMiddleware,额外功能:
|
||||
- 保留原有的 summary 历史记录逻辑
|
||||
- 将 ToolMessage 的 results 卸载到虚拟文件系统
|
||||
1. 触发 Summary 时,卸载超过阈值的工具结果
|
||||
2. 智能保留策略:
|
||||
- 触发 Summary 时,首先进行卸载
|
||||
- 只有当总 Token 数超过 max_retention_ratio * trigger 时,才进行消息清理(Summary)
|
||||
- 始终保留 System Message
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: str | BaseChatModel,
|
||||
*,
|
||||
trigger: ContextSize | list[ContextSize] | None = None,
|
||||
keep: ContextSize = ("messages", _DEFAULT_MESSAGES_TO_KEEP),
|
||||
token_counter: TokenCounter = count_tokens_approximately,
|
||||
summary_prompt: str = DEFAULT_SUMMARY_PROMPT,
|
||||
trim_tokens_to_summarize: int | None = 4000,
|
||||
# 工具结果卸载参数
|
||||
summary_offload_threshold: int = 1000,
|
||||
max_retention_ratio: float = 0.6,
|
||||
**deprecated_kwargs: Any,
|
||||
) -> None:
|
||||
"""初始化中间件.
|
||||
|
||||
Args:
|
||||
model: 用于生成摘要的语言模型
|
||||
trigger: 触发摘要的阈值条件 (建议使用 ("tokens", N))
|
||||
keep: 摘要后保留的消息数量/ token 策略 (作为 fallback)
|
||||
token_counter: token 计数函数
|
||||
summary_prompt: 生成摘要的提示词模板
|
||||
trim_tokens_to_summarize: Summary 时,无损保留的消息数
|
||||
summary_offload_threshold: Summary 时,工具调用结果超过此 token 数阈值则卸载到文件系统
|
||||
max_retention_ratio: 触发 Summary 后,如果不超过此比例(相对于 trigger),则不删除消息。默认 0.6
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
if isinstance(model, str):
|
||||
model = init_chat_model(model)
|
||||
|
||||
self.model = model
|
||||
if trigger is None:
|
||||
self.trigger: ContextSize | list[ContextSize] | None = None
|
||||
trigger_conditions: list[ContextSize] = []
|
||||
elif isinstance(trigger, list):
|
||||
validated_list = [self._validate_context_size(item, "trigger") for item in trigger]
|
||||
self.trigger = validated_list
|
||||
trigger_conditions = validated_list
|
||||
else:
|
||||
validated = self._validate_context_size(trigger, "trigger")
|
||||
self.trigger = validated
|
||||
trigger_conditions = [validated]
|
||||
self._trigger_conditions = trigger_conditions
|
||||
|
||||
self.keep = self._validate_context_size(keep, "keep")
|
||||
if token_counter is count_tokens_approximately:
|
||||
self.token_counter = _get_approximate_token_counter(self.model)
|
||||
else:
|
||||
self.token_counter = token_counter
|
||||
self.summary_prompt = summary_prompt
|
||||
self.trim_tokens_to_summarize = trim_tokens_to_summarize
|
||||
|
||||
# 工具结果卸载配置
|
||||
self.summary_offload_threshold = summary_offload_threshold
|
||||
self.max_retention_ratio = max_retention_ratio
|
||||
|
||||
# 检查 fractional 配置需要的 model profile
|
||||
requires_profile = any(condition[0] == "fraction" for condition in self._trigger_conditions)
|
||||
if self.keep[0] == "fraction":
|
||||
requires_profile = True
|
||||
if requires_profile and self._get_profile_limits() is None:
|
||||
msg = (
|
||||
"Model profile information is required to use fractional token limits, "
|
||||
"and is unavailable for the specified model. Please use absolute token "
|
||||
"counts instead, or pass "
|
||||
'`ChatModel(..., profile={"max_input_tokens": ...})`.'
|
||||
)
|
||||
raise ValueError(msg)
|
||||
|
||||
def _get_token_trigger_value(self) -> int | None:
|
||||
"""Helper to get the token trigger value."""
|
||||
if not self._trigger_conditions:
|
||||
return None
|
||||
|
||||
for kind, value in self._trigger_conditions:
|
||||
if kind == "tokens":
|
||||
return int(value)
|
||||
# Support fractional if needed, converting to tokens using profile
|
||||
if kind == "fraction":
|
||||
max_input_tokens = self._get_profile_limits()
|
||||
if max_input_tokens:
|
||||
return int(max_input_tokens * value)
|
||||
return None
|
||||
|
||||
@override
|
||||
def before_model(self, state: AgentState[Any], runtime: Runtime) -> dict[str, Any] | None:
|
||||
"""Process messages before model invocation, potentially triggering summarization."""
|
||||
|
||||
messages = state["messages"]
|
||||
|
||||
self._ensure_message_ids(messages)
|
||||
|
||||
total_tokens = self.token_counter(messages)
|
||||
|
||||
# 1. 检查是否触发 Summary
|
||||
if not self._should_summarize(messages, total_tokens):
|
||||
return None
|
||||
|
||||
# 2. 触发 Summary:卸载超阈值的工具结果
|
||||
files_update: dict[str, Any] = {}
|
||||
modified_messages: list[AnyMessage] = []
|
||||
|
||||
agg_files, agg_msgs = _offload_tool_results(
|
||||
messages, self.summary_offload_threshold, self.token_counter, runtime
|
||||
)
|
||||
files_update = agg_files
|
||||
modified_messages = agg_msgs
|
||||
|
||||
# 3. 检查 Retention Ratio
|
||||
current_tokens = self.token_counter(messages)
|
||||
trigger_value = self._get_token_trigger_value()
|
||||
|
||||
if trigger_value is None:
|
||||
system_msg_count = 1 if messages and messages[0].type == "system" else 0
|
||||
messages_to_process = messages[1:] if system_msg_count else messages
|
||||
cutoff_relative = self._determine_cutoff_index(messages_to_process)
|
||||
cutoff_index = system_msg_count + cutoff_relative
|
||||
else:
|
||||
retention_limit = trigger_value * self.max_retention_ratio
|
||||
|
||||
if current_tokens <= retention_limit:
|
||||
if files_update:
|
||||
result: dict[str, Any] = {"messages": modified_messages}
|
||||
if files_update:
|
||||
result["files"] = files_update
|
||||
return result
|
||||
return None
|
||||
|
||||
# 4. 超过 limit,需要 Eviction (Summary)
|
||||
system_msg_count = 0
|
||||
messages_to_process = messages
|
||||
|
||||
if messages and messages[0].type == "system":
|
||||
system_msg_count = 1
|
||||
messages_to_process = messages[1:]
|
||||
|
||||
cutoff_relative = self._find_cutoff_by_token_limit(messages_to_process, int(retention_limit))
|
||||
cutoff_index = system_msg_count + cutoff_relative
|
||||
|
||||
if cutoff_index <= system_msg_count:
|
||||
if files_update:
|
||||
result = {"messages": modified_messages}
|
||||
if files_update:
|
||||
result["files"] = files_update
|
||||
return result
|
||||
return None
|
||||
|
||||
system_message = messages[0] if messages and messages[0].type == "system" else None
|
||||
conversation_messages = messages[1:] if system_message is not None else messages
|
||||
|
||||
messages_to_summarize, preserved_messages = self._partition_messages(
|
||||
conversation_messages, cutoff_index - system_msg_count
|
||||
)
|
||||
summary = self._create_summary(messages_to_summarize)
|
||||
new_messages = self._build_new_messages(summary)
|
||||
|
||||
# 如果有 System Message,需要保留在最前面
|
||||
final_messages = []
|
||||
|
||||
if system_message is not None:
|
||||
final_messages.append(system_message)
|
||||
|
||||
final_messages.extend(new_messages)
|
||||
final_messages.extend(preserved_messages)
|
||||
|
||||
result: dict[str, Any] = {"messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES), *final_messages]}
|
||||
|
||||
if files_update:
|
||||
result["files"] = files_update
|
||||
|
||||
return result
|
||||
|
||||
@override
|
||||
async def abefore_model(self, state: AgentState[Any], runtime: Runtime) -> dict[str, Any] | None:
|
||||
"""Process messages before model invocation, potentially triggering summarization."""
|
||||
|
||||
messages = state["messages"]
|
||||
|
||||
self._ensure_message_ids(messages)
|
||||
|
||||
total_tokens = self.token_counter(messages)
|
||||
|
||||
# 1. 检查是否触发 Summary
|
||||
if not self._should_summarize(messages, total_tokens):
|
||||
return None
|
||||
|
||||
# 2. 触发 Summary:卸载超阈值的工具结果
|
||||
files_update: dict[str, Any] = {}
|
||||
modified_messages: list[AnyMessage] = []
|
||||
|
||||
agg_files, agg_msgs = await _aoffload_tool_results(
|
||||
messages, self.summary_offload_threshold, self.token_counter, runtime
|
||||
)
|
||||
files_update = agg_files
|
||||
modified_messages = agg_msgs
|
||||
|
||||
# 3. 检查 Retention Ratio
|
||||
current_tokens = self.token_counter(messages)
|
||||
trigger_value = self._get_token_trigger_value()
|
||||
|
||||
if trigger_value is None:
|
||||
system_msg_count = 1 if messages and messages[0].type == "system" else 0
|
||||
messages_to_process = messages[1:] if system_msg_count else messages
|
||||
cutoff_relative = self._determine_cutoff_index(messages_to_process)
|
||||
cutoff_index = system_msg_count + cutoff_relative
|
||||
else:
|
||||
retention_limit = trigger_value * self.max_retention_ratio
|
||||
|
||||
if current_tokens <= retention_limit:
|
||||
if files_update:
|
||||
result: dict[str, Any] = {"messages": modified_messages}
|
||||
if files_update:
|
||||
result["files"] = files_update
|
||||
return result
|
||||
return None
|
||||
|
||||
# 4. 超过 limit,需要 Eviction (Summary)
|
||||
system_msg_count = 0
|
||||
messages_to_process = messages
|
||||
|
||||
if messages and messages[0].type == "system":
|
||||
system_msg_count = 1
|
||||
messages_to_process = messages[1:]
|
||||
|
||||
cutoff_relative = self._find_cutoff_by_token_limit(messages_to_process, int(retention_limit))
|
||||
cutoff_index = system_msg_count + cutoff_relative
|
||||
|
||||
if cutoff_index <= system_msg_count:
|
||||
if files_update:
|
||||
result = {"messages": modified_messages}
|
||||
if files_update:
|
||||
result["files"] = files_update
|
||||
return result
|
||||
return None
|
||||
|
||||
system_message = messages[0] if messages and messages[0].type == "system" else None
|
||||
conversation_messages = messages[1:] if system_message is not None else messages
|
||||
|
||||
messages_to_summarize, preserved_messages = self._partition_messages(
|
||||
conversation_messages, cutoff_index - system_msg_count
|
||||
)
|
||||
|
||||
summary = await self._acreate_summary(messages_to_summarize)
|
||||
new_messages = self._build_new_messages(summary)
|
||||
|
||||
final_messages = []
|
||||
|
||||
if system_message is not None:
|
||||
final_messages.append(system_message)
|
||||
|
||||
final_messages.extend(new_messages)
|
||||
final_messages.extend(preserved_messages)
|
||||
|
||||
result: dict[str, Any] = {"messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES), *final_messages]}
|
||||
|
||||
if files_update:
|
||||
result["files"] = files_update
|
||||
|
||||
return result
|
||||
|
||||
def _should_summarize(self, messages: list[AnyMessage], total_tokens: int) -> bool:
|
||||
"""Determine whether summarization should run for the current token usage."""
|
||||
if not self._trigger_conditions:
|
||||
return False
|
||||
|
||||
for kind, value in self._trigger_conditions:
|
||||
if kind == "messages" and len(messages) >= value:
|
||||
return True
|
||||
if kind == "tokens" and total_tokens >= value:
|
||||
return True
|
||||
if kind == "fraction":
|
||||
max_input_tokens = self._get_profile_limits()
|
||||
if max_input_tokens is None:
|
||||
continue
|
||||
threshold = int(max_input_tokens * value)
|
||||
if threshold <= 0:
|
||||
threshold = 1
|
||||
if total_tokens >= threshold:
|
||||
return True
|
||||
return False
|
||||
|
||||
def _determine_cutoff_index(self, messages: list[AnyMessage]) -> int:
|
||||
"""Choose cutoff index respecting retention configuration."""
|
||||
kind, value = self.keep
|
||||
if kind in {"tokens", "fraction"}:
|
||||
token_based_cutoff = self._find_token_based_cutoff(messages)
|
||||
if token_based_cutoff is not None:
|
||||
return token_based_cutoff
|
||||
return self._find_safe_cutoff(messages, _DEFAULT_MESSAGES_TO_KEEP)
|
||||
return self._find_safe_cutoff(messages, cast("int", value))
|
||||
|
||||
def _find_token_based_cutoff(self, messages: list[AnyMessage]) -> int | None:
|
||||
"""Find cutoff index based on target token retention."""
|
||||
if not messages:
|
||||
return 0
|
||||
|
||||
kind, value = self.keep
|
||||
if kind == "fraction":
|
||||
max_input_tokens = self._get_profile_limits()
|
||||
if max_input_tokens is None:
|
||||
return None
|
||||
target_token_count = int(max_input_tokens * value)
|
||||
elif kind == "tokens":
|
||||
target_token_count = int(value)
|
||||
else:
|
||||
return None
|
||||
|
||||
if target_token_count <= 0:
|
||||
target_token_count = 1
|
||||
|
||||
if self.token_counter(messages) <= target_token_count:
|
||||
return 0
|
||||
|
||||
# 二分查找
|
||||
left, right = 0, len(messages)
|
||||
cutoff_candidate = len(messages)
|
||||
max_iterations = len(messages).bit_length() + 1
|
||||
for _ in range(max_iterations):
|
||||
if left >= right:
|
||||
break
|
||||
|
||||
mid = (left + right) // 2
|
||||
if self.token_counter(messages[mid:]) <= target_token_count:
|
||||
cutoff_candidate = mid
|
||||
right = mid
|
||||
else:
|
||||
left = mid + 1
|
||||
|
||||
if cutoff_candidate == len(messages):
|
||||
cutoff_candidate = left
|
||||
|
||||
if cutoff_candidate >= len(messages):
|
||||
if len(messages) == 1:
|
||||
return 0
|
||||
cutoff_candidate = len(messages) - 1
|
||||
|
||||
return self._find_safe_cutoff_point(messages, cutoff_candidate)
|
||||
|
||||
def _find_cutoff_by_token_limit(self, messages: list[AnyMessage], max_tokens: int) -> int:
|
||||
"""Find cutoff index to ensure total tokens <= max_tokens."""
|
||||
if not messages or self.token_counter(messages) <= max_tokens:
|
||||
return 0
|
||||
|
||||
# Binary search for cutoff
|
||||
left, right = 0, len(messages)
|
||||
cutoff_candidate = len(messages)
|
||||
max_iterations = len(messages).bit_length() + 1
|
||||
|
||||
for _ in range(max_iterations):
|
||||
if left >= right:
|
||||
break
|
||||
|
||||
mid = (left + right) // 2
|
||||
# Calculate tokens for preserved part: messages[mid:]
|
||||
if self.token_counter(messages[mid:]) <= max_tokens:
|
||||
cutoff_candidate = mid
|
||||
right = mid
|
||||
else:
|
||||
left = mid + 1
|
||||
|
||||
if cutoff_candidate == len(messages):
|
||||
cutoff_candidate = left
|
||||
|
||||
return self._find_safe_cutoff_point(messages, cutoff_candidate)
|
||||
|
||||
def _get_profile_limits(self) -> int | None:
|
||||
"""Retrieve max input token limit from the model profile."""
|
||||
try:
|
||||
profile = self.model.profile
|
||||
except AttributeError:
|
||||
return None
|
||||
|
||||
if not isinstance(profile, Mapping):
|
||||
return None
|
||||
|
||||
max_input_tokens = profile.get("max_input_tokens")
|
||||
|
||||
if not isinstance(max_input_tokens, int):
|
||||
return None
|
||||
|
||||
return max_input_tokens
|
||||
|
||||
@staticmethod
|
||||
def _validate_context_size(context: ContextSize, parameter_name: str) -> ContextSize:
|
||||
"""Validate context configuration tuples."""
|
||||
kind, value = context
|
||||
if kind == "fraction":
|
||||
if not 0 < value <= 1:
|
||||
msg = f"Fractional {parameter_name} values must be between 0 and 1, got {value}."
|
||||
raise ValueError(msg)
|
||||
elif kind in {"tokens", "messages"}:
|
||||
if value <= 0:
|
||||
msg = f"{parameter_name} thresholds must be greater than 0, got {value}."
|
||||
raise ValueError(msg)
|
||||
else:
|
||||
msg = f"Unsupported context size type {kind} for {parameter_name}."
|
||||
raise ValueError(msg)
|
||||
return context
|
||||
|
||||
@staticmethod
|
||||
def _build_new_messages(summary: str) -> list[HumanMessage]:
|
||||
return [
|
||||
HumanMessage(
|
||||
content=f"Here is a summary of the conversation to date:\n\n{summary}",
|
||||
additional_kwargs={"lc_source": "summarization"},
|
||||
)
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _ensure_message_ids(messages: list[AnyMessage]) -> None:
|
||||
"""Ensure all messages have unique IDs for the add_messages reducer."""
|
||||
for msg in messages:
|
||||
if msg.id is None:
|
||||
msg.id = str(uuid.uuid4())
|
||||
|
||||
@staticmethod
|
||||
def _partition_messages(
|
||||
conversation_messages: list[AnyMessage],
|
||||
cutoff_index: int,
|
||||
) -> tuple[list[AnyMessage], list[AnyMessage]]:
|
||||
"""Partition messages into those to summarize and those to preserve."""
|
||||
messages_to_summarize = conversation_messages[:cutoff_index]
|
||||
preserved_messages = conversation_messages[cutoff_index:]
|
||||
|
||||
return messages_to_summarize, preserved_messages
|
||||
|
||||
def _find_safe_cutoff(self, messages: list[AnyMessage], messages_to_keep: int) -> int:
|
||||
"""Find safe cutoff point that preserves AI/Tool message pairs."""
|
||||
if len(messages) <= messages_to_keep:
|
||||
return 0
|
||||
|
||||
target_cutoff = len(messages) - messages_to_keep
|
||||
return self._find_safe_cutoff_point(messages, target_cutoff)
|
||||
|
||||
@staticmethod
|
||||
def _find_safe_cutoff_point(messages: list[AnyMessage], cutoff_index: int) -> int:
|
||||
"""Find a safe cutoff point that doesn't split AI/Tool message pairs."""
|
||||
if cutoff_index >= len(messages) or not isinstance(messages[cutoff_index], ToolMessage):
|
||||
return cutoff_index
|
||||
|
||||
tool_call_ids: set[str] = set()
|
||||
idx = cutoff_index
|
||||
while idx < len(messages) and isinstance(messages[idx], ToolMessage):
|
||||
tool_msg = cast("ToolMessage", messages[idx])
|
||||
if tool_msg.tool_call_id:
|
||||
tool_call_ids.add(tool_msg.tool_call_id)
|
||||
idx += 1
|
||||
|
||||
for i in range(cutoff_index - 1, -1, -1):
|
||||
msg = messages[i]
|
||||
if isinstance(msg, AIMessage) and msg.tool_calls:
|
||||
ai_tool_call_ids = {tc.get("id") for tc in msg.tool_calls if tc.get("id")}
|
||||
if tool_call_ids & ai_tool_call_ids:
|
||||
return i
|
||||
|
||||
return idx
|
||||
|
||||
def _create_summary(self, messages_to_summarize: list[AnyMessage]) -> str:
|
||||
"""Generate summary for the given messages."""
|
||||
if not messages_to_summarize:
|
||||
return "No previous conversation history."
|
||||
|
||||
trimmed_messages = self._trim_messages_for_summary(messages_to_summarize)
|
||||
if not trimmed_messages:
|
||||
return "Previous conversation was too long to summarize."
|
||||
|
||||
formatted_messages = get_buffer_string(trimmed_messages)
|
||||
|
||||
try:
|
||||
response = self.model.invoke(self.summary_prompt.format(messages=formatted_messages))
|
||||
return response.text.strip()
|
||||
except Exception as e:
|
||||
return f"Error generating summary: {e!s}"
|
||||
|
||||
async def _acreate_summary(self, messages_to_summarize: list[AnyMessage]) -> str:
|
||||
"""Generate summary for the given messages."""
|
||||
if not messages_to_summarize:
|
||||
return "No previous conversation history."
|
||||
|
||||
trimmed_messages = self._trim_messages_for_summary(messages_to_summarize)
|
||||
if not trimmed_messages:
|
||||
return "Previous conversation was too long to summarize."
|
||||
|
||||
formatted_messages = get_buffer_string(trimmed_messages)
|
||||
|
||||
try:
|
||||
response = await self.model.ainvoke(self.summary_prompt.format(messages=formatted_messages))
|
||||
return response.text.strip()
|
||||
except Exception as e:
|
||||
return f"Error generating summary: {e!s}"
|
||||
|
||||
def _trim_messages_for_summary(self, messages: list[AnyMessage]) -> list[AnyMessage]:
|
||||
"""Trim messages to fit within summary generation limits."""
|
||||
try:
|
||||
if self.trim_tokens_to_summarize is None:
|
||||
return messages
|
||||
return cast(
|
||||
"list[AnyMessage]",
|
||||
trim_messages(
|
||||
messages,
|
||||
max_tokens=self.trim_tokens_to_summarize,
|
||||
token_counter=self.token_counter,
|
||||
start_on="human",
|
||||
strategy="last",
|
||||
allow_partial=True,
|
||||
include_system=True,
|
||||
),
|
||||
)
|
||||
except Exception:
|
||||
return messages[-_DEFAULT_FALLBACK_MESSAGE_COUNT:]
|
||||
|
||||
|
||||
# 便捷函数:创建中间件实例
|
||||
def create_summary_offload_middleware(
|
||||
def create_summary_middleware(
|
||||
model: str | BaseChatModel,
|
||||
*,
|
||||
trigger: ContextSize | list[ContextSize] | None = None,
|
||||
keep: ContextSize = ("messages", _DEFAULT_MESSAGES_TO_KEEP),
|
||||
summary_offload_threshold: int = 1000,
|
||||
max_retention_ratio: float = 0.6,
|
||||
) -> SummaryOffloadMiddleware:
|
||||
"""创建 SummaryOffloadMiddleware 实例的便捷函数"""
|
||||
return SummaryOffloadMiddleware(
|
||||
trigger: ContextSize | list[ContextSize] | None,
|
||||
keep: ContextSize,
|
||||
trim_tokens_to_summarize: int | None = 4000,
|
||||
) -> SummarizationMiddleware:
|
||||
"""Create DeepAgents summarization middleware using Yuxi's virtual outputs root."""
|
||||
middleware = SummarizationMiddleware(
|
||||
model=model,
|
||||
backend=create_agent_composite_backend,
|
||||
trigger=trigger,
|
||||
keep=keep,
|
||||
summary_offload_threshold=summary_offload_threshold,
|
||||
max_retention_ratio=max_retention_ratio,
|
||||
trim_tokens_to_summarize=trim_tokens_to_summarize,
|
||||
)
|
||||
middleware._history_path_prefix = VIRTUAL_PATH_CONVERSATION_HISTORY
|
||||
middleware._large_tool_results_prefix = VIRTUAL_PATH_LARGE_TOOL_RESULTS
|
||||
return middleware
|
||||
|
||||
@ -5,13 +5,14 @@ from contextlib import asynccontextmanager
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
from deepagents.middleware.subagents import GENERAL_PURPOSE_SUBAGENT
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from yuxi.agents.subagents.repository import SubAgentRepository
|
||||
from yuxi.storage.postgres.manager import pg_manager
|
||||
from yuxi.storage.postgres.models_business import SubAgent
|
||||
from yuxi.utils import logger
|
||||
from yuxi.utils.paths import OUTPUTS_DIR_NAME
|
||||
from yuxi.utils.paths import VIRTUAL_PATH_OUTPUTS
|
||||
|
||||
# SubAgent specs cache for get_subagent_specs
|
||||
_subagent_specs_cache: list[dict[str, Any]] | None = None
|
||||
@ -38,7 +39,7 @@ _DEFAULT_SUBAGENTS = [
|
||||
"你是一位专注的研究员。你的工作是根据用户的问题进行研究。"
|
||||
"进行彻底的研究,然后用详细的答案回复用户的问题,只有你的最终答案会被传递给用户。"
|
||||
"除了你的最终信息,他们不会知道任何其他事情,所以你的最终报告应该就是你的最终信息!"
|
||||
f"将调研结果保存到主题研究文件中 {OUTPUTS_DIR_NAME}/sub_research/xxx.md 中。"
|
||||
f"将调研结果保存到主题研究文件中 {VIRTUAL_PATH_OUTPUTS}/sub_research/xxx.md 中。"
|
||||
),
|
||||
"tools": ["tavily_search"],
|
||||
"is_builtin": True,
|
||||
@ -134,6 +135,40 @@ def clear_specs_cache() -> None:
|
||||
_subagent_specs_cache = None
|
||||
|
||||
|
||||
def build_subagent_middleware_specs(
|
||||
subagents: list[dict[str, Any]],
|
||||
*,
|
||||
default_model: Any,
|
||||
default_tools: list[Any] | None = None,
|
||||
default_middleware: list[Any] | None = None,
|
||||
include_general_purpose: bool = True,
|
||||
model_loader: Any | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
tools = list(default_tools or [])
|
||||
middleware = list(default_middleware or [])
|
||||
specs: list[dict[str, Any]] = []
|
||||
|
||||
if include_general_purpose:
|
||||
specs.append(
|
||||
{
|
||||
**GENERAL_PURPOSE_SUBAGENT,
|
||||
"model": default_model,
|
||||
"tools": list(tools),
|
||||
"middleware": list(middleware),
|
||||
}
|
||||
)
|
||||
|
||||
for spec in subagents:
|
||||
item = {key: value for key, value in spec.items() if key != "slug"}
|
||||
model_name = item.get("model")
|
||||
item["model"] = model_loader(model_name) if model_name and model_loader else (model_name or default_model)
|
||||
item["tools"] = list(item.get("tools") or tools)
|
||||
item["middleware"] = [*middleware, *list(item.get("middleware", []))]
|
||||
specs.append(item)
|
||||
|
||||
return specs
|
||||
|
||||
|
||||
async def get_subagents_from_slugs(selected_slugs: Any, *, db: AsyncSession | None = None) -> list[dict[str, Any]]:
|
||||
"""根据 slug 获取 subagent specs(含工具解析)。"""
|
||||
specs = await get_subagent_specs(db)
|
||||
|
||||
@ -46,7 +46,10 @@
|
||||
|
||||
|
||||
# def _collect_sandbox_file_paths(backend, remote_dir: str) -> list[str]:
|
||||
# entries = backend.ls_info(remote_dir)
|
||||
# result = backend.ls(remote_dir)
|
||||
# if result.error:
|
||||
# raise ValueError(result.error)
|
||||
# entries = result.entries or []
|
||||
# file_paths: list[str] = []
|
||||
# for entry in entries:
|
||||
# path = entry["path"]
|
||||
|
||||
@ -38,6 +38,7 @@ from yuxi.utils.paths import VIRTUAL_PATH_OUTPUTS, VIRTUAL_PATH_UPLOADS, VIRTUAL
|
||||
|
||||
_PROTECTED_USER_DATA_ROOTS = frozenset(
|
||||
{
|
||||
USER_DATA_PATH,
|
||||
VIRTUAL_PATH_WORKSPACE,
|
||||
VIRTUAL_PATH_UPLOADS,
|
||||
VIRTUAL_PATH_OUTPUTS,
|
||||
@ -281,8 +282,10 @@ async def list_viewer_filesystem_tree(
|
||||
return {"entries": _sort_entries(entries)}
|
||||
|
||||
if _is_skills_path(normalized_path):
|
||||
entries = await asyncio.to_thread(skills_backend.ls_info, _strip_skills_prefix(normalized_path))
|
||||
remapped = [_remap_prefixed_entry(entry, SKILLS_PATH) for entry in entries]
|
||||
result = await asyncio.to_thread(skills_backend.ls, _strip_skills_prefix(normalized_path))
|
||||
if result.error:
|
||||
raise HTTPException(status_code=400, detail=result.error)
|
||||
remapped = [_remap_prefixed_entry(entry, SKILLS_PATH) for entry in (result.entries or [])]
|
||||
return {"entries": _sort_entries(remapped)}
|
||||
except PermissionError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e)) from e
|
||||
|
||||
@ -8,11 +8,15 @@ WORKSPACE_AGENTS_DIR_NAME = "agents"
|
||||
WORKSPACE_AGENTS_PROMPT_FILE_NAME = "AGENTS.md"
|
||||
UPLOADS_DIR_NAME = "uploads"
|
||||
OUTPUTS_DIR_NAME = "outputs"
|
||||
LARGE_TOOL_RESULTS_DIR_NAME = "large_tool_results"
|
||||
CONVERSATION_HISTORY_DIR_NAME = "conversation_history"
|
||||
VIRTUAL_SKILLS_PATH = "/home/gem/skills"
|
||||
|
||||
VIRTUAL_PATH_WORKSPACE = (Path(VIRTUAL_PATH_PREFIX) / WORKSPACE_DIR_NAME).as_posix()
|
||||
VIRTUAL_PATH_UPLOADS = (Path(VIRTUAL_PATH_PREFIX) / UPLOADS_DIR_NAME).as_posix()
|
||||
VIRTUAL_PATH_OUTPUTS = (Path(VIRTUAL_PATH_PREFIX) / OUTPUTS_DIR_NAME).as_posix()
|
||||
VIRTUAL_PATH_LARGE_TOOL_RESULTS = (Path(VIRTUAL_PATH_OUTPUTS) / LARGE_TOOL_RESULTS_DIR_NAME).as_posix()
|
||||
VIRTUAL_PATH_CONVERSATION_HISTORY = (Path(VIRTUAL_PATH_OUTPUTS) / CONVERSATION_HISTORY_DIR_NAME).as_posix()
|
||||
|
||||
__all__ = [
|
||||
"VIRTUAL_PATH_PREFIX",
|
||||
@ -21,8 +25,12 @@ __all__ = [
|
||||
"WORKSPACE_AGENTS_PROMPT_FILE_NAME",
|
||||
"UPLOADS_DIR_NAME",
|
||||
"OUTPUTS_DIR_NAME",
|
||||
"LARGE_TOOL_RESULTS_DIR_NAME",
|
||||
"CONVERSATION_HISTORY_DIR_NAME",
|
||||
"VIRTUAL_PATH_WORKSPACE",
|
||||
"VIRTUAL_PATH_UPLOADS",
|
||||
"VIRTUAL_PATH_OUTPUTS",
|
||||
"VIRTUAL_PATH_LARGE_TOOL_RESULTS",
|
||||
"VIRTUAL_PATH_CONVERSATION_HISTORY",
|
||||
"VIRTUAL_SKILLS_PATH",
|
||||
]
|
||||
|
||||
@ -4,7 +4,9 @@ from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import uuid
|
||||
from types import SimpleNamespace
|
||||
|
||||
from fastapi import HTTPException
|
||||
import pytest
|
||||
from yuxi.agents.backends.sandbox import (
|
||||
ensure_thread_dirs,
|
||||
@ -95,8 +97,8 @@ async def test_viewer_tree_root_does_not_require_sandbox_listing(test_client, st
|
||||
thread_id = await _create_thread_for_user(test_client, headers)
|
||||
|
||||
class _FailingSandbox:
|
||||
def ls_info(self, path):
|
||||
raise AssertionError(f"sandbox ls_info should not be used for root path: {path}")
|
||||
def ls(self, path):
|
||||
raise AssertionError(f"sandbox ls should not be used for root path: {path}")
|
||||
|
||||
class _EmptyBackend:
|
||||
def has_entries(self):
|
||||
@ -131,8 +133,8 @@ async def test_viewer_tree_user_data_uses_local_thread_directory(test_client, st
|
||||
actual_path.write_text("viewer tree", encoding="utf-8")
|
||||
|
||||
class _FailingSandbox:
|
||||
def ls_info(self, path):
|
||||
raise AssertionError(f"sandbox ls_info should not be used for user-data path: {path}")
|
||||
def ls(self, path):
|
||||
raise AssertionError(f"sandbox ls should not be used for user-data path: {path}")
|
||||
|
||||
class _EmptyBackend:
|
||||
def has_entries(self):
|
||||
@ -448,6 +450,7 @@ async def test_viewer_delete_rejects_readonly_namespace_directory(test_client, s
|
||||
@pytest.mark.parametrize(
|
||||
"protected_path",
|
||||
[
|
||||
"/home/gem/user-data",
|
||||
"/home/gem/user-data/workspace",
|
||||
"/home/gem/user-data/uploads",
|
||||
"/home/gem/user-data/outputs",
|
||||
@ -471,6 +474,30 @@ async def test_viewer_delete_rejects_protected_user_data_root_directories(
|
||||
assert response.json()["detail"] == "当前目录不允许删除"
|
||||
|
||||
|
||||
async def test_delete_viewer_file_rejects_user_data_root_without_removing(monkeypatch):
|
||||
from yuxi.services import viewer_filesystem_service as service_module
|
||||
|
||||
async def _fake_resolve_viewer_state(**kwargs):
|
||||
return object(), object(), []
|
||||
|
||||
def _fail_rmtree(path):
|
||||
raise AssertionError(f"unexpected root removal: {path}")
|
||||
|
||||
monkeypatch.setattr(service_module, "_resolve_viewer_state", _fake_resolve_viewer_state)
|
||||
monkeypatch.setattr(service_module.shutil, "rmtree", _fail_rmtree)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await service_module.delete_viewer_file(
|
||||
thread_id="thread-1",
|
||||
path="/home/gem/user-data",
|
||||
current_user=SimpleNamespace(uid="user-1"),
|
||||
db=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert exc_info.value.detail == "当前目录不允许删除"
|
||||
|
||||
|
||||
async def test_viewer_create_directory_adds_workspace_folder(test_client, standard_user):
|
||||
headers = standard_user["headers"]
|
||||
uid = str(standard_user["user"]["uid"])
|
||||
|
||||
@ -9,15 +9,12 @@ pytestmark = [pytest.mark.asyncio, pytest.mark.integration]
|
||||
|
||||
|
||||
async def _create_thread_for_user(test_client, headers: dict[str, str]) -> str:
|
||||
agents_resp = await test_client.get("/api/chat/agent", headers=headers)
|
||||
assert agents_resp.status_code == 200, agents_resp.text
|
||||
agents = agents_resp.json().get("agents", [])
|
||||
if not agents:
|
||||
pytest.skip("No agents available for viewer filesystem integration tests.")
|
||||
|
||||
agent_id = agents[0].get("id")
|
||||
agent_resp = await test_client.get("/api/agent/default", headers=headers)
|
||||
assert agent_resp.status_code == 200, agent_resp.text
|
||||
agent = agent_resp.json().get("agent") or {}
|
||||
agent_id = agent.get("slug") or agent.get("id")
|
||||
if not agent_id:
|
||||
pytest.skip("Agent payload missing id field.")
|
||||
pytest.skip("Default agent payload missing id field.")
|
||||
|
||||
create_resp = await test_client.post(
|
||||
"/api/chat/thread",
|
||||
|
||||
@ -7,6 +7,7 @@ from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from deepagents.backends.protocol import LsResult
|
||||
from langgraph.types import Command
|
||||
|
||||
from yuxi.agents.toolkits.buildin.install_skill import (
|
||||
@ -112,13 +113,15 @@ def test_download_skill_dir_preserves_nested_files(tmp_path):
|
||||
calls = []
|
||||
|
||||
class Backend:
|
||||
def ls_info(self, path):
|
||||
def ls(self, path):
|
||||
if path.endswith("/demo"):
|
||||
return [
|
||||
{"path": f"{path}/SKILL.md", "is_dir": False},
|
||||
{"path": f"{path}/scripts", "is_dir": True},
|
||||
]
|
||||
return [{"path": f"{path}/run.py", "is_dir": False}]
|
||||
return LsResult(
|
||||
entries=[
|
||||
{"path": f"{path}/SKILL.md", "is_dir": False},
|
||||
{"path": f"{path}/scripts", "is_dir": True},
|
||||
]
|
||||
)
|
||||
return LsResult(entries=[{"path": f"{path}/run.py", "is_dir": False}])
|
||||
|
||||
def download_files(self, paths):
|
||||
calls.append(paths)
|
||||
@ -142,8 +145,8 @@ def test_download_skill_dir_raises_on_download_error(tmp_path):
|
||||
"""_download_skill_dir should fail instead of importing a partial skill."""
|
||||
|
||||
class Backend:
|
||||
def ls_info(self, path):
|
||||
return [{"path": f"{path}/SKILL.md", "is_dir": False}]
|
||||
def ls(self, path):
|
||||
return LsResult(entries=[{"path": f"{path}/SKILL.md", "is_dir": False}])
|
||||
|
||||
def download_files(self, paths):
|
||||
return [SimpleNamespace(path=paths[0], content=None, error="file_not_found")]
|
||||
|
||||
@ -5,10 +5,16 @@ from __future__ import annotations
|
||||
from types import MethodType, SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from yuxi.agents.backends.composite import create_agent_composite_backend
|
||||
from deepagents.backends.protocol import GlobResult
|
||||
from yuxi.agents.backends.composite import (
|
||||
CustomCompositeBackend,
|
||||
create_agent_composite_backend,
|
||||
create_agent_filesystem_middleware,
|
||||
)
|
||||
from yuxi.agents.backends.sandbox import resolve_virtual_path, sandbox_id_for_thread
|
||||
from yuxi.agents.backends.sandbox.backend import ProvisionerSandboxBackend
|
||||
from yuxi.agents.middlewares.skills_middleware import SkillsMiddleware
|
||||
from yuxi.utils.paths import VIRTUAL_PATH_CONVERSATION_HISTORY, VIRTUAL_PATH_LARGE_TOOL_RESULTS
|
||||
|
||||
|
||||
def _runtime(
|
||||
@ -40,6 +46,7 @@ def test_create_agent_composite_backend_uses_prepared_readable_skills(monkeypatc
|
||||
|
||||
assert isinstance(backend.default, ProvisionerSandboxBackend)
|
||||
assert backend.default._readable_skills == ["reporter"]
|
||||
assert backend.artifacts_root == "/home/gem/user-data/outputs"
|
||||
assert "/skills/" in backend.routes
|
||||
assert "/home/gem/kbs/" not in backend.routes
|
||||
|
||||
@ -57,6 +64,35 @@ def test_create_agent_composite_backend_ignores_unprepared_context_skills(monkey
|
||||
assert backend.default._readable_skills == []
|
||||
|
||||
|
||||
def test_create_agent_filesystem_middleware_uses_outputs_for_internal_artifacts() -> None:
|
||||
middleware = create_agent_filesystem_middleware(tool_token_limit_before_evict=500)
|
||||
|
||||
assert middleware._tool_token_limit_before_evict == 500
|
||||
assert middleware._large_tool_results_prefix == VIRTUAL_PATH_LARGE_TOOL_RESULTS
|
||||
assert middleware._conversation_history_prefix == VIRTUAL_PATH_CONVERSATION_HISTORY
|
||||
|
||||
|
||||
def test_custom_composite_glob_only_searches_routes_from_root() -> None:
|
||||
class _Backend:
|
||||
def __init__(self, name: str):
|
||||
self.name = name
|
||||
self.calls: list[tuple[str, str]] = []
|
||||
|
||||
def glob(self, pattern: str, path: str = "/") -> GlobResult:
|
||||
self.calls.append((pattern, path))
|
||||
return GlobResult(matches=[{"path": f"{path.rstrip('/')}/{self.name}.md"}])
|
||||
|
||||
default = _Backend("default")
|
||||
routed = _Backend("skill")
|
||||
backend = CustomCompositeBackend(default=default, routes={"/skills/": routed})
|
||||
|
||||
result = backend.glob("**/*.md", path="/home/gem/user-data")
|
||||
|
||||
assert result.error is None
|
||||
assert default.calls == [("**/*.md", "/home/gem/user-data")]
|
||||
assert routed.calls == []
|
||||
|
||||
|
||||
def test_skills_middleware_extracts_slug_for_new_paths() -> None:
|
||||
middleware = SkillsMiddleware()
|
||||
assert middleware.skills_sources_for_prompt == ["/home/gem/skills/"]
|
||||
@ -82,6 +118,73 @@ def test_sandbox_id_for_thread_is_stable():
|
||||
assert len(sid1) == 12
|
||||
|
||||
|
||||
def test_provisioner_denies_reads_outside_allowed_roots(monkeypatch) -> None:
|
||||
monkeypatch.setattr("yuxi.agents.backends.sandbox.backend.get_sandbox_provider", lambda: object())
|
||||
backend = ProvisionerSandboxBackend(thread_id="thread-1", uid="user-1")
|
||||
|
||||
result = backend.read("/etc/passwd")
|
||||
|
||||
assert result.error == "permission denied for read on '/etc/passwd'"
|
||||
|
||||
|
||||
def test_provisioner_denies_upload_writes(monkeypatch) -> None:
|
||||
monkeypatch.setattr("yuxi.agents.backends.sandbox.backend.get_sandbox_provider", lambda: object())
|
||||
backend = ProvisionerSandboxBackend(thread_id="thread-1", uid="user-1")
|
||||
|
||||
write_result = backend.write("/home/gem/user-data/uploads/blocked.txt", "blocked")
|
||||
upload_result = backend.upload_files([("/home/gem/user-data/uploads/blocked.bin", b"blocked")])
|
||||
|
||||
assert write_result.error and "permission denied" in write_result.error
|
||||
assert upload_result[0].error == "permission_denied"
|
||||
|
||||
|
||||
def test_provisioner_allows_outputs_writes(monkeypatch) -> None:
|
||||
monkeypatch.setattr("yuxi.agents.backends.sandbox.backend.get_sandbox_provider", lambda: object())
|
||||
backend = ProvisionerSandboxBackend(thread_id="thread-1", uid="user-1")
|
||||
|
||||
def _missing_file(path, offset=0, limit=None):
|
||||
raise FileNotFoundError
|
||||
|
||||
monkeypatch.setattr(backend, "_read_binary", _missing_file)
|
||||
|
||||
calls = []
|
||||
|
||||
def _write_file(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return SimpleNamespace(success=True, message="")
|
||||
|
||||
fake_client = SimpleNamespace(file=SimpleNamespace(write_file=_write_file))
|
||||
backend._get_client = MethodType(lambda self: fake_client, backend)
|
||||
|
||||
result = backend.write("/home/gem/user-data/outputs/report.md", "ok")
|
||||
|
||||
assert result.error is None
|
||||
assert result.path == "/home/gem/user-data/outputs/report.md"
|
||||
assert calls[0]["file"] == "/home/gem/user-data/outputs/report.md"
|
||||
|
||||
|
||||
def test_provisioner_glob_root_searches_readable_roots(monkeypatch) -> None:
|
||||
monkeypatch.setattr("yuxi.agents.backends.sandbox.backend.get_sandbox_provider", lambda: object())
|
||||
backend = ProvisionerSandboxBackend(thread_id="thread-1", uid="user-1")
|
||||
calls = []
|
||||
|
||||
def _find_files(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return SimpleNamespace(data=SimpleNamespace(files=[f"{kwargs['path']}/match.md"]))
|
||||
|
||||
fake_client = SimpleNamespace(file=SimpleNamespace(find_files=_find_files))
|
||||
backend._get_client = MethodType(lambda self: fake_client, backend)
|
||||
|
||||
result = backend.glob("**/*.md")
|
||||
|
||||
assert result.error is None
|
||||
assert [call["path"] for call in calls] == ["/home/gem/user-data", "/home/gem/skills"]
|
||||
assert [item["path"] for item in result.matches] == [
|
||||
"/home/gem/skills/match.md",
|
||||
"/home/gem/user-data/match.md",
|
||||
]
|
||||
|
||||
|
||||
def test_provisioner_read_reports_binary_files(monkeypatch) -> None:
|
||||
monkeypatch.setattr("yuxi.agents.backends.sandbox.backend.get_sandbox_provider", lambda: object())
|
||||
backend = ProvisionerSandboxBackend(thread_id="thread-1", uid="user-1")
|
||||
@ -89,21 +192,27 @@ def test_provisioner_read_reports_binary_files(monkeypatch) -> None:
|
||||
|
||||
result = backend.read("/home/gem/user-data/image.png")
|
||||
|
||||
assert result == "Error: File '/home/gem/user-data/image.png' is binary and cannot be rendered as text"
|
||||
assert result.error is None
|
||||
assert result.file_data is not None
|
||||
assert result.file_data["encoding"] == "base64"
|
||||
|
||||
|
||||
def test_provisioner_read_reports_invalid_path(monkeypatch) -> None:
|
||||
monkeypatch.setattr("yuxi.agents.backends.sandbox.backend.get_sandbox_provider", lambda: object())
|
||||
backend = ProvisionerSandboxBackend(thread_id="thread-1", uid="user-1")
|
||||
|
||||
def _raise_invalid_path(path, offset=0, limit=None):
|
||||
raise ValueError("path traversal is not allowed")
|
||||
result = backend.read("secret.txt")
|
||||
|
||||
monkeypatch.setattr(backend, "_read_binary", _raise_invalid_path)
|
||||
assert result.error == "Invalid path 'secret.txt': path must start with /"
|
||||
|
||||
result = backend.read("../secret.txt")
|
||||
|
||||
assert result == "Error: Invalid path '../secret.txt': path traversal is not allowed"
|
||||
def test_provisioner_read_reports_path_traversal(monkeypatch) -> None:
|
||||
monkeypatch.setattr("yuxi.agents.backends.sandbox.backend.get_sandbox_provider", lambda: object())
|
||||
backend = ProvisionerSandboxBackend(thread_id="thread-1", uid="user-1")
|
||||
|
||||
result = backend.read("/home/gem/user-data/../secret.txt")
|
||||
|
||||
assert result.error == "Invalid path '/home/gem/user-data/../secret.txt': path traversal is not allowed"
|
||||
|
||||
|
||||
def test_provisioner_download_files_distinguishes_invalid_path_from_read_failure(monkeypatch) -> None:
|
||||
@ -111,13 +220,11 @@ def test_provisioner_download_files_distinguishes_invalid_path_from_read_failure
|
||||
backend = ProvisionerSandboxBackend(thread_id="thread-1", uid="user-1")
|
||||
|
||||
def _fake_read_binary(path, offset=0, limit=None):
|
||||
if path == "/bad-path":
|
||||
raise ValueError("path is required")
|
||||
raise RuntimeError("sandbox read timeout")
|
||||
|
||||
monkeypatch.setattr(backend, "_read_binary", _fake_read_binary)
|
||||
|
||||
responses = backend.download_files(["/bad-path", "/read-failed"])
|
||||
responses = backend.download_files(["bad-path", "/home/gem/user-data/read-failed"])
|
||||
|
||||
assert responses[0].error == "invalid_path"
|
||||
assert responses[1].error.startswith("read_failed")
|
||||
|
||||
@ -24,8 +24,8 @@ def test_selected_skills_backend_none_exposes_no_skills(tmp_path, monkeypatch):
|
||||
|
||||
backend = skills_backend.SelectedSkillsReadonlyBackend(selected_slugs=None)
|
||||
|
||||
assert backend.ls_info("/") == []
|
||||
assert "Access denied" in backend.read("/alpha/SKILL.md")
|
||||
assert backend.ls("/").entries == []
|
||||
assert backend.read("/alpha/SKILL.md").error == "Access denied: file is outside selected skills."
|
||||
|
||||
|
||||
def test_selected_skills_backend_readonly_and_visible_only_selected(tmp_path, monkeypatch):
|
||||
@ -34,15 +34,17 @@ def test_selected_skills_backend_readonly_and_visible_only_selected(tmp_path, mo
|
||||
|
||||
backend = skills_backend.SelectedSkillsReadonlyBackend(selected_slugs=["alpha"])
|
||||
|
||||
root_entries = backend.ls_info("/")
|
||||
paths = sorted(entry.get("path") for entry in root_entries)
|
||||
root_result = backend.ls("/")
|
||||
assert root_result.error is None
|
||||
paths = sorted(entry.get("path") for entry in (root_result.entries or []))
|
||||
assert paths == ["/alpha/"]
|
||||
|
||||
ok_read = backend.read("/alpha/SKILL.md")
|
||||
assert "alpha" in ok_read
|
||||
assert ok_read.file_data is not None
|
||||
assert "alpha" in ok_read.file_data["content"]
|
||||
|
||||
denied_read = backend.read("/beta/SKILL.md")
|
||||
assert "Access denied" in denied_read
|
||||
assert denied_read.error and "Access denied" in denied_read.error
|
||||
|
||||
write_result = backend.write("/alpha/new.md", "x")
|
||||
assert write_result.error and "read-only" in write_result.error
|
||||
|
||||
@ -9,8 +9,7 @@ from yuxi.agents.backends.skills_backend import SelectedSkillsReadonlyBackend
|
||||
def test_skills_backend_read_outside_root_returns_error_message(monkeypatch) -> None:
|
||||
def _fake_read(self, file_path: str, offset: int = 0, limit: int = 2000):
|
||||
raise ValueError(
|
||||
"Path:/app/package/yuxi/agents/skills/buildin/reporter "
|
||||
"outside root directory: /app/saves/skills"
|
||||
"Path:/app/package/yuxi/agents/skills/buildin/reporter outside root directory: /app/saves/skills"
|
||||
)
|
||||
|
||||
monkeypatch.setattr(FilesystemBackend, "read", _fake_read)
|
||||
@ -20,12 +19,12 @@ def test_skills_backend_read_outside_root_returns_error_message(monkeypatch) ->
|
||||
backend.read("/reporter/SKILL.md")
|
||||
|
||||
|
||||
def test_skills_backend_ls_info_outside_root_returns_empty(monkeypatch) -> None:
|
||||
def _fake_ls_info(self, path: str):
|
||||
def test_skills_backend_ls_outside_root_raises(monkeypatch) -> None:
|
||||
def _fake_ls(self, path: str):
|
||||
raise ValueError("Path outside root directory")
|
||||
|
||||
monkeypatch.setattr(FilesystemBackend, "ls_info", _fake_ls_info)
|
||||
monkeypatch.setattr(FilesystemBackend, "ls", _fake_ls)
|
||||
|
||||
backend = SelectedSkillsReadonlyBackend(selected_slugs=["reporter"])
|
||||
with pytest.raises(ValueError, match="outside root directory"):
|
||||
backend.ls_info("/reporter")
|
||||
backend.ls("/reporter")
|
||||
|
||||
@ -3,13 +3,17 @@ from __future__ import annotations
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import ToolMessage
|
||||
from langchain_core.messages import SystemMessage, ToolMessage
|
||||
from langgraph.types import Command
|
||||
|
||||
import yuxi.agents.middlewares.skills_middleware as skills_middleware
|
||||
from yuxi.agents.middlewares.skills_middleware import SkillsMiddleware, resolve_runtime_skills_for_context
|
||||
|
||||
|
||||
def _system_message_text(message: SystemMessage) -> str:
|
||||
return "\n".join(block.get("text", "") for block in message.content_blocks if isinstance(block, dict))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_runtime_skills_derives_prompt_and_readable_closure(monkeypatch):
|
||||
async def fake_list_skills_from_db(db=None, user=None):
|
||||
@ -47,9 +51,9 @@ async def test_resolve_runtime_skills_derives_prompt_and_readable_closure(monkey
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skills_prompt_uses_prepared_prompt_skills():
|
||||
async def test_skills_prompt_uses_prepared_prompt_skills_at_request_level():
|
||||
context = SimpleNamespace(
|
||||
system_prompt="base",
|
||||
system_prompt="context base",
|
||||
skills=["configured-only"],
|
||||
_prompt_skills=["alpha"],
|
||||
_runtime_skill_metadata={
|
||||
@ -66,12 +70,34 @@ async def test_skills_prompt_uses_prepared_prompt_skills():
|
||||
},
|
||||
)
|
||||
|
||||
await SkillsMiddleware().abefore_agent({}, SimpleNamespace(context=context))
|
||||
class FakeRequest:
|
||||
def __init__(self, *, system_message=None, tools=None):
|
||||
self.runtime = SimpleNamespace(context=context)
|
||||
self.state = {}
|
||||
self.tools = tools or []
|
||||
self.system_message = system_message or SystemMessage(content="base")
|
||||
|
||||
assert "base" in context.system_prompt
|
||||
assert "Alpha" in context.system_prompt
|
||||
assert "Configured Only" not in context.system_prompt
|
||||
assert getattr(context, "_skills_prompt_injected") is True
|
||||
def override(self, **kwargs):
|
||||
return FakeRequest(
|
||||
system_message=kwargs.get("system_message", self.system_message),
|
||||
tools=kwargs.get("tools", self.tools),
|
||||
)
|
||||
|
||||
captured = {}
|
||||
|
||||
async def handler(request):
|
||||
captured["system_message"] = request.system_message
|
||||
return "ok"
|
||||
|
||||
result = await SkillsMiddleware().awrap_model_call(FakeRequest(), handler)
|
||||
prompt_text = _system_message_text(captured["system_message"])
|
||||
|
||||
assert result == "ok"
|
||||
assert "base" in prompt_text
|
||||
assert "Alpha" in prompt_text
|
||||
assert "Configured Only" not in prompt_text
|
||||
assert context.system_prompt == "context base"
|
||||
assert not hasattr(context, "_skills_prompt_injected")
|
||||
assert not hasattr(context, "_visible_skills")
|
||||
|
||||
|
||||
|
||||
@ -3,11 +3,11 @@ from __future__ import annotations
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage
|
||||
from deepagents.middleware.summarization import SummarizationMiddleware
|
||||
|
||||
import yuxi.agents.middlewares.summary_middleware as summary_middleware
|
||||
from yuxi.agents.middlewares.summary_middleware import SummaryOffloadMiddleware
|
||||
from yuxi.utils.paths import VIRTUAL_PATH_OUTPUTS
|
||||
from yuxi.agents.backends.composite import create_agent_composite_backend
|
||||
from yuxi.agents.middlewares.summary_middleware import create_summary_middleware
|
||||
from yuxi.utils.paths import VIRTUAL_PATH_CONVERSATION_HISTORY, VIRTUAL_PATH_LARGE_TOOL_RESULTS
|
||||
|
||||
|
||||
class _DummyModel:
|
||||
@ -19,120 +19,18 @@ class _DummyModel:
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_offload_tool_result_writes_readable_outputs_path(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
def _fake_write(_runtime, file_path: str, content: str) -> tuple[bool, dict]:
|
||||
captured["file_path"] = file_path
|
||||
captured["content"] = content
|
||||
return True, {}
|
||||
|
||||
monkeypatch.setattr(summary_middleware, "_write_offloaded_content", _fake_write)
|
||||
|
||||
message = ToolMessage(content="line1\nline2", tool_call_id="tool-1", name="search", id="msg-1")
|
||||
result = summary_middleware._offload_tool_result(
|
||||
message,
|
||||
threshold=1,
|
||||
token_counter=lambda _messages: 10,
|
||||
runtime=SimpleNamespace(),
|
||||
)
|
||||
|
||||
assert result == {}
|
||||
assert captured["file_path"] == f"{VIRTUAL_PATH_OUTPUTS}/summary_offload/search-msg-1.txt"
|
||||
assert captured["content"].startswith("=== Tool Invocation ===\nTool: search\nTool Call ID: tool-1\n")
|
||||
assert "文件路径: " in str(message.content)
|
||||
assert captured["file_path"] in str(message.content)
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_offload_tool_result_skips_non_text_content(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
called = False
|
||||
|
||||
def _fake_write(_runtime, _file_path: str, _content: str) -> tuple[bool, dict]:
|
||||
nonlocal called
|
||||
called = True
|
||||
return True, {}
|
||||
|
||||
monkeypatch.setattr(summary_middleware, "_write_offloaded_content", _fake_write)
|
||||
|
||||
message = ToolMessage(
|
||||
content=[{"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "abc"}}],
|
||||
tool_call_id="tool-1",
|
||||
name="read_file",
|
||||
id="msg-2",
|
||||
)
|
||||
|
||||
result = summary_middleware._offload_tool_result(
|
||||
message,
|
||||
threshold=1,
|
||||
token_counter=lambda _messages: 10,
|
||||
runtime=SimpleNamespace(),
|
||||
)
|
||||
|
||||
assert result is None
|
||||
assert called is False
|
||||
assert message.content[0]["type"] == "image"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_before_model_excludes_system_message_from_summary(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
middleware = SummaryOffloadMiddleware(
|
||||
def test_create_summary_middleware_uses_deepagents_with_yuxi_outputs_root() -> None:
|
||||
middleware = create_summary_middleware(
|
||||
model=_DummyModel(),
|
||||
trigger=("tokens", 10),
|
||||
keep=("messages", 1),
|
||||
token_counter=lambda _messages: 100,
|
||||
summary_offload_threshold=10_000,
|
||||
max_retention_ratio=0.5,
|
||||
)
|
||||
captured_ids: list[str] = []
|
||||
|
||||
monkeypatch.setattr(summary_middleware, "_offload_tool_results", lambda *args, **kwargs: ({}, []))
|
||||
monkeypatch.setattr(middleware, "_find_cutoff_by_token_limit", lambda _messages, _limit: 2)
|
||||
monkeypatch.setattr(
|
||||
middleware,
|
||||
"_create_summary",
|
||||
lambda messages: captured_ids.extend([str(message.id) for message in messages]) or "summary",
|
||||
trigger=("tokens", 90_000),
|
||||
keep=("tokens", 45_000),
|
||||
trim_tokens_to_summarize=4000,
|
||||
)
|
||||
|
||||
messages = [
|
||||
SystemMessage(content="sys", id="sys-1"),
|
||||
HumanMessage(content="human-1", id="human-1"),
|
||||
AIMessage(content="ai-1", id="ai-1"),
|
||||
HumanMessage(content="human-2", id="human-2"),
|
||||
]
|
||||
|
||||
result = middleware.before_model({"messages": messages}, SimpleNamespace())
|
||||
|
||||
assert captured_ids == ["human-1", "ai-1"]
|
||||
assert result is not None
|
||||
new_messages = result["messages"]
|
||||
assert new_messages[1].id == "sys-1"
|
||||
assert new_messages[2].content == "Here is a summary of the conversation to date:\n\nsummary"
|
||||
assert new_messages[3].id == "human-2"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_before_model_uses_keep_cutoff_for_message_trigger(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
middleware = SummaryOffloadMiddleware(
|
||||
model=_DummyModel(),
|
||||
trigger=("messages", 3),
|
||||
keep=("messages", 1),
|
||||
token_counter=lambda _messages: 1,
|
||||
summary_offload_threshold=10_000,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(summary_middleware, "_offload_tool_results", lambda *args, **kwargs: ({}, []))
|
||||
monkeypatch.setattr(middleware, "_create_summary", lambda _messages: "summary")
|
||||
|
||||
messages = [
|
||||
HumanMessage(content="human-1", id="human-1"),
|
||||
AIMessage(content="ai-1", id="ai-1"),
|
||||
HumanMessage(content="human-2", id="human-2"),
|
||||
]
|
||||
|
||||
result = middleware.before_model({"messages": messages}, SimpleNamespace())
|
||||
|
||||
assert result is not None
|
||||
new_messages = result["messages"]
|
||||
assert new_messages[1].content == "Here is a summary of the conversation to date:\n\nsummary"
|
||||
assert new_messages[2].id == "human-2"
|
||||
assert isinstance(middleware, SummarizationMiddleware)
|
||||
assert middleware._backend is create_agent_composite_backend
|
||||
assert middleware._history_path_prefix == VIRTUAL_PATH_CONVERSATION_HISTORY
|
||||
assert middleware._large_tool_results_prefix == VIRTUAL_PATH_LARGE_TOOL_RESULTS
|
||||
assert middleware._lc_helper.trigger == ("tokens", 90_000)
|
||||
assert middleware._lc_helper.keep == ("tokens", 45_000)
|
||||
assert middleware._lc_helper.trim_tokens_to_summarize == 4000
|
||||
|
||||
@ -642,6 +642,35 @@ class TestDeepAgentSubagentSelection:
|
||||
assert [item["name"] for item in resolved_specs] == ["research-agent"]
|
||||
assert resolved_specs[0]["tools"] == [mock_tool]
|
||||
|
||||
def test_builtin_research_subagent_uses_virtual_outputs_path(self):
|
||||
from yuxi.agents.subagents import service as service_module
|
||||
|
||||
research_agent = next(item for item in service_module._DEFAULT_SUBAGENTS if item["slug"] == "research-agent")
|
||||
|
||||
assert "/home/gem/user-data/outputs/sub_research/xxx.md" in research_agent["system_prompt"]
|
||||
|
||||
def test_build_subagent_middleware_specs_uses_default_tools_for_empty_tools(self):
|
||||
from yuxi.agents.subagents import service as service_module
|
||||
|
||||
default_tool = MagicMock()
|
||||
|
||||
specs = service_module.build_subagent_middleware_specs(
|
||||
[
|
||||
{
|
||||
"slug": "empty-tools-agent",
|
||||
"name": "empty-tools-agent",
|
||||
"description": "",
|
||||
"system_prompt": "s",
|
||||
"tools": [],
|
||||
}
|
||||
],
|
||||
default_model="test-model",
|
||||
default_tools=[default_tool],
|
||||
include_general_purpose=False,
|
||||
)
|
||||
|
||||
assert specs[0]["tools"] == [default_tool]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_subagents_from_slugs_none_selects_none(self, monkeypatch):
|
||||
from yuxi.agents.subagents import service as service_module
|
||||
|
||||
2537
backend/uv.lock
2537
backend/uv.lock
File diff suppressed because it is too large
Load Diff
@ -48,7 +48,7 @@
|
||||
- 新增 Milvus 图谱检索链路:Query 可召回图谱实体和三元组,结合 Chunk 命中实体构造 seed entity,读取 Neo4j 2-hop 子图后用 igraph 执行 PPR,最终以 Chunk 为产物并通过 RRF 与原 Chunk 召回融合;检索配置改为 dataclass 元数据生成,支持 `depend_on` 控制重排序和图检索参数展示。
|
||||
- 收紧用户管理部门隔离:普通管理员创建用户时固定归属本部门,用户列表、访问选项、详情、更新和删除接口均限制在本部门范围内。
|
||||
- 调整 Agent 资源默认选择与运行时上下文:未显式配置工具、知识库、MCP、Skills、SubAgent 时默认启用当前用户可访问/可用的全部资源,显式保存空列表仍表示不启用对应资源;Agent 创建前统一完成最终资源权限过滤、知识库 `db_id` 可见范围派生和 Skill prompt/readable 依赖闭包派生,聊天运行时与文件系统预览复用同一结果。
|
||||
- 重构 Skills 权限与安装流程:Skill 增加 `source_type/share_config/enabled`,内置 Skill 作为启动同步入库的全局资源,不再保留前端安装/更新状态,支持启停但不允许删除;上传和远程添加统一为解析草稿后确认生效范围,管理端支持编辑生效范围与启停;Agent 运行时按当前用户可访问 Skills 派生 prompt/readable 依赖闭包并限制挂载/激活。
|
||||
- 重构 Skills 权限与安装流程:Skill 增加 `source_type/share_config/enabled`,内置 Skill 作为启动同步入库的全局资源,不再保留前端安装/更新状态,支持启停但不允许删除;上传和远程添加统一为解析草稿后确认生效范围,管理端支持编辑生效范围与启停;Agent 运行时按当前用户可访问 Skills 派生 prompt/readable 依赖闭包并限制挂载/激活,Skills prompt 改为模型请求级注入以避免污染 runtime context。
|
||||
- 精简历史兼容层:移除 sandbox provisioner `local` 后端别名、ask_user_question 单问题旧协议、JWT 历史默认密钥特殊判断、内置 Skill `SKILLS.md` 文件名回退、运行事件数字 seq 兼容和前端若干旧字段回退。
|
||||
- 重构知识库共享权限:`share_config` 改为全局共享、部门共享、指定人可访问三档,部门共享必须包含当前用户部门,指定人可访问必须包含当前用户,并补充权限过滤测试。
|
||||
- 移除知识库沙盒文件系统映射:不再通过 `/home/gem/kbs` 暴露知识库文件树,Agent 继续使用 `query_kb` 与 `open_kb_document` 访问知识库内容。
|
||||
@ -70,6 +70,7 @@
|
||||
- 标准化 Agent run/SSE 执行链路:run 创建时持久化输入消息并提交后入队,worker 统一写入 Redis Stream envelope,SSE 输出 `event/data/id`、心跳注释、`Last-Event-ID` 回放和终止 `end` 事件;前端强制使用 run API 并支持 ask_user_question 中断后以 resume run 恢复。
|
||||
- 收敛后端模块边界:文档解析从 `plugins.parser` 移动到 `knowledge.parser`,内容审查从 `plugins.guard` 移动到 `services.guard`。
|
||||
- 收敛文件服务边界:文件预览判断抽为独立服务,Viewer 文件系统的 workspace 分支复用用户 workspace 服务,线程运行时上下文解析从泛化 `filesystem_service` 拆出为 agent runtime helper。
|
||||
- 升级 DeepAgents 到 0.6.7 并适配新版文件系统协议:SubAgentMiddleware 改为显式 subagent spec,Skills prompt 补齐新版占位符;sandbox/skills backend 复用新版 `ReadResult`、`GlobResult`、`GrepResult` 等协议类型,文件权限在 backend 层明确区分 skills、uploads、outputs 与 workspace,保留最小 `CustomCompositeBackend` 以避免非 route glob 误扫其他 route;Agent 上下文压缩改为复用 DeepAgents SummarizationMiddleware,历史摘要与大工具结果统一 offload 到 outputs。
|
||||
- 优化聊天输入 @ 文件提及:未创建 Thread 时可搜索用户 workspace,创建 Thread 后按当前对话文件优先、workspace 兜底的来源顺序搜索,并拆分 workspace/thread 缓存避免假 thread 与跨用户缓存污染;输入框与用户消息支持将 raw mention 渲染为带类型图标的引用单元,文件仅显示文件名且保留原始沙盒路径文本。
|
||||
|
||||
---
|
||||
|
||||
@ -1942,8 +1942,8 @@ watch(currentChatId, (threadId, oldThreadId) => {
|
||||
border: 1px solid var(--gray-150);
|
||||
border-radius: 16px;
|
||||
box-shadow:
|
||||
0 16px 40px rgba(15, 23, 42, 0.1),
|
||||
0 2px 10px rgba(15, 23, 42, 0.06);
|
||||
0 16px 40px var(--shadow-1),
|
||||
0 2px 10px var(--shadow-0);
|
||||
z-index: 20;
|
||||
margin: 0 8px 8px;
|
||||
margin-left: -16px;
|
||||
@ -2151,7 +2151,7 @@ watch(currentChatId, (threadId, oldThreadId) => {
|
||||
max-width: 800px;
|
||||
margin: 0 auto;
|
||||
flex-grow: 1;
|
||||
padding: 1rem 1.5rem;
|
||||
padding: 1rem var(--page-padding);
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
}
|
||||
|
||||
@ -760,7 +760,7 @@ watch(
|
||||
min-height: 0;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
padding: 0 6px;
|
||||
padding: 6px;
|
||||
}
|
||||
|
||||
.files-display.has-preview.with-tree .tree-pane {
|
||||
|
||||
@ -646,12 +646,22 @@ const checkMentionTrigger = () => {
|
||||
|
||||
// 更新提及候选项
|
||||
const updateMentionItems = (query = '') => {
|
||||
const normalizedQuery = String(query || '')
|
||||
if (!normalizedQuery) {
|
||||
clearTimeout(mentionSearchTimer)
|
||||
if (activeAbortController) {
|
||||
activeAbortController.abort()
|
||||
activeAbortController = null
|
||||
}
|
||||
searchRequestId.value++
|
||||
}
|
||||
|
||||
if (!props.mention) {
|
||||
mentionItems.value = { files: [], knowledgeBases: [], mcps: [], skills: [], subagents: [] }
|
||||
return
|
||||
}
|
||||
|
||||
const lowerQuery = query.toLowerCase()
|
||||
const lowerQuery = normalizedQuery.toLowerCase()
|
||||
const { files = [], knowledgeBases = [], mcps = [], skills = [], subagents = [] } = props.mention
|
||||
|
||||
const filterItems = (list) =>
|
||||
@ -686,7 +696,7 @@ const updateMentionItems = (query = '') => {
|
||||
}
|
||||
})
|
||||
|
||||
const filteredLocalFiles = query ? filterItems(localFileItems) : []
|
||||
const filteredLocalFiles = normalizedQuery ? filterItems(localFileItems) : []
|
||||
|
||||
const knowledgeItems = knowledgeBases.map((kb) => {
|
||||
const kbName = kb.name || ''
|
||||
@ -749,23 +759,23 @@ const updateMentionItems = (query = '') => {
|
||||
subagents: filterItems(subagentItems)
|
||||
}
|
||||
|
||||
if (query) {
|
||||
if (normalizedQuery) {
|
||||
const activeThreadId = props.threadId || ''
|
||||
clearTimeout(mentionSearchTimer)
|
||||
mentionSearchTimer = setTimeout(async () => {
|
||||
// 物理中断之前的未完成 HTTP 请求
|
||||
if (activeAbortController) {
|
||||
activeAbortController.abort()
|
||||
}
|
||||
activeAbortController = new AbortController()
|
||||
if (activeAbortController) {
|
||||
activeAbortController.abort()
|
||||
activeAbortController = null
|
||||
}
|
||||
searchRequestId.value++
|
||||
const currentId = searchRequestId.value
|
||||
|
||||
searchRequestId.value++
|
||||
const currentId = searchRequestId.value
|
||||
mentionSearchTimer = setTimeout(async () => {
|
||||
activeAbortController = new AbortController()
|
||||
|
||||
try {
|
||||
const responseData = await searchMentionFiles(
|
||||
activeThreadId,
|
||||
query,
|
||||
normalizedQuery,
|
||||
activeAbortController.signal
|
||||
)
|
||||
|
||||
|
||||
@ -9,11 +9,18 @@ export const mentionTypePrefixMap = {
|
||||
}
|
||||
|
||||
const mentionTypePattern = Object.values(mentionTypePrefixMap).join('|')
|
||||
const mentionTokenRegex = new RegExp(`@(${mentionTypePattern}):\\S+`, 'g')
|
||||
const mentionTokenRegex = new RegExp(`@(${mentionTypePattern}):(?:"((?:\\\\.|[^"\\\\])*)"|(\\S+))`, 'g')
|
||||
|
||||
const quoteMentionValue = (value) => String(value ?? '').replace(/\\/g, '\\\\').replace(/"/g, '\\"')
|
||||
const unquoteMentionValue = (value) => String(value ?? '').replace(/\\(["\\])/g, '$1')
|
||||
|
||||
export const formatMentionToken = (type, value) => {
|
||||
const prefix = mentionTypePrefixMap[type] || type
|
||||
return `@${prefix}:${value}`
|
||||
const rawValue = String(value ?? '')
|
||||
if (/\s|["\\]/.test(rawValue)) {
|
||||
return `@${prefix}:"${quoteMentionValue(rawValue)}"`
|
||||
}
|
||||
return `@${prefix}:${rawValue}`
|
||||
}
|
||||
|
||||
export const parseMentionText = (text = '') => {
|
||||
@ -37,12 +44,13 @@ export const parseMentionText = (text = '') => {
|
||||
}
|
||||
|
||||
const type = match[1]
|
||||
const prefix = `@${type}:`
|
||||
const quotedValue = match[2]
|
||||
const rawValue = match[3]
|
||||
segments.push({
|
||||
kind: 'mention',
|
||||
raw,
|
||||
type,
|
||||
value: raw.slice(prefix.length),
|
||||
value: quotedValue !== undefined ? unquoteMentionValue(quotedValue) : rawValue,
|
||||
start,
|
||||
end
|
||||
})
|
||||
@ -63,7 +71,7 @@ export const parseMentionText = (text = '') => {
|
||||
}
|
||||
|
||||
const setMentionLabel = (labels, type, value, label) => {
|
||||
const rawValue = String(value || '').trim()
|
||||
const rawValue = String(value ?? '').trim()
|
||||
const displayLabel = String(label || '').trim()
|
||||
if (!rawValue || !displayLabel) return
|
||||
labels[`${type}:${rawValue}`] = displayLabel
|
||||
@ -112,10 +120,10 @@ export const getMentionDisplayLabel = (type, value, displayLabels = {}) => {
|
||||
if (mappedLabel) return mappedLabel
|
||||
|
||||
if (type === 'file') {
|
||||
const normalizedPath = String(value || '').replace(/\/+$/, '')
|
||||
const normalizedPath = String(value ?? '').replace(/\/+$/, '')
|
||||
return getDisplayFileName(normalizedPath || value, '文件')
|
||||
}
|
||||
return String(value || '').trim() || type
|
||||
return String(value ?? '').trim() || type
|
||||
}
|
||||
|
||||
export const findActiveMentionQuery = (text = '', rawCaretOffset = 0) => {
|
||||
|
||||
Loading…
Reference in New Issue
Block a user