feat: 适配 deepagents 最新版特性

This commit is contained in:
Wenjie Zhang 2026-05-31 13:36:42 +08:00
parent b51df7097e
commit 31a7fc52d0
29 changed files with 1964 additions and 2475 deletions

View File

@ -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",

View File

@ -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",

View File

@ -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

View File

@ -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:

View File

@ -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.")

View File

@ -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)

View File

@ -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="任务结束前,应该检查维护的待办事项列表是否结束。"),

View File

@ -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",
]

View File

@ -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,
)

View File

@ -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

View File

@ -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)

View File

@ -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"]

View File

@ -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

View File

@ -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",
]

View File

@ -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"])

View File

@ -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",

View File

@ -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")]

View File

@ -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")

View File

@ -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

View File

@ -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")

View File

@ -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")

View File

@ -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

View File

@ -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

File diff suppressed because it is too large Load Diff

View File

@ -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 envelopeSSE 输出 `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 specSkills prompt 补齐新版占位符sandbox/skills backend 复用新版 `ReadResult`、`GlobResult`、`GrepResult` 等协议类型,文件权限在 backend 层明确区分 skills、uploads、outputs 与 workspace保留最小 `CustomCompositeBackend` 以避免非 route glob 误扫其他 routeAgent 上下文压缩改为复用 DeepAgents SummarizationMiddleware历史摘要与大工具结果统一 offload 到 outputs。
- 优化聊天输入 @ 文件提及:未创建 Thread 时可搜索用户 workspace创建 Thread 后按当前对话文件优先、workspace 兜底的来源顺序搜索,并拆分 workspace/thread 缓存避免假 thread 与跨用户缓存污染;输入框与用户消息支持将 raw mention 渲染为带类型图标的引用单元,文件仅显示文件名且保留原始沙盒路径文本。
---

View File

@ -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;
}

View File

@ -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 {

View File

@ -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
)

View File

@ -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) => {