fix: resolve ruff async path violations

This commit is contained in:
LIghtJUNction
2026-03-18 17:36:46 +08:00
parent c9fac2bf82
commit fb93949edd
5 changed files with 70 additions and 49 deletions
+12 -5
View File
@@ -1246,7 +1246,7 @@ class PluginManager:
_, repo_name, _ = self.updator.parse_github_url(repo_url)
repo_name = self.updator.format_name(repo_name)
plugin_path = os.path.join(self.plugin_store_path, repo_name)
if os.path.exists(plugin_path):
if await anyio.Path(plugin_path).exists():
raise Exception(
f"安装失败:目录 {os.path.basename(plugin_path)} 已存在。"
)
@@ -1259,8 +1259,9 @@ class PluginManager:
self.plugin_store_path,
metadata_dir_name,
)
if target_plugin_path != plugin_path and os.path.exists(
target_plugin_path
if (
target_plugin_path != plugin_path
and await anyio.Path(target_plugin_path).exists()
):
raise Exception(f"安装失败:目录 {metadata_dir_name} 已存在。")
if target_plugin_path != plugin_path:
@@ -1649,7 +1650,10 @@ class PluginManager:
self.plugin_store_path,
metadata_dir_name,
)
if target_plugin_path != desti_dir and os.path.exists(target_plugin_path):
if (
target_plugin_path != desti_dir
and await anyio.Path(target_plugin_path).exists()
):
raise Exception(f"安装失败:目录 {metadata_dir_name} 已存在。")
if target_plugin_path != desti_dir:
os.rename(desti_dir, target_plugin_path)
@@ -1724,7 +1728,10 @@ class PluginManager:
)
raise
finally:
if temp_desti_dir != desti_dir and os.path.isdir(temp_desti_dir):
if (
temp_desti_dir != desti_dir
and await anyio.Path(temp_desti_dir).is_dir()
):
try:
remove_dir(temp_desti_dir)
except Exception as e:
+24 -14
View File
@@ -4,8 +4,10 @@ import os
import re
import uuid
from contextlib import asynccontextmanager
from pathlib import Path
from typing import cast
import anyio
from quart import Response as QuartResponse
from quart import g, make_response, request, send_file
@@ -50,6 +52,10 @@ async def _poll_webchat_stream_result(back_queue, username: str):
return result, False
def _resolve_path(path: str) -> Path:
return Path(path).resolve(strict=False)
class ChatRoute(Route):
def __init__(
self,
@@ -95,27 +101,29 @@ class ChatRoute(Route):
try:
file_path = os.path.join(self.attachments_dir, os.path.basename(filename))
real_file_path = os.path.realpath(file_path)
real_imgs_dir = os.path.realpath(self.attachments_dir)
resolved_file_path = _resolve_path(file_path)
resolved_base_dir = _resolve_path(self.attachments_dir)
if not os.path.exists(real_file_path):
if not await anyio.Path(resolved_file_path).exists():
# try legacy
file_path = os.path.join(
self.legacy_img_dir, os.path.basename(filename)
)
if os.path.exists(file_path):
real_file_path = os.path.realpath(file_path)
real_imgs_dir = os.path.realpath(self.legacy_img_dir)
if await anyio.Path(file_path).exists():
resolved_file_path = _resolve_path(file_path)
resolved_base_dir = _resolve_path(self.legacy_img_dir)
if not real_file_path.startswith(real_imgs_dir):
try:
resolved_file_path.relative_to(resolved_base_dir)
except ValueError:
return Response().error("Invalid file path").__dict__
filename_ext = os.path.splitext(filename)[1].lower()
if filename_ext == ".wav":
return await send_file(real_file_path, mimetype="audio/wav")
return await send_file(str(resolved_file_path), mimetype="audio/wav")
if filename_ext[1:] in self.supported_imgs:
return await send_file(real_file_path, mimetype="image/jpeg")
return await send_file(real_file_path)
return await send_file(str(resolved_file_path), mimetype="image/jpeg")
return await send_file(str(resolved_file_path))
except (FileNotFoundError, OSError):
return Response().error("File access error").__dict__
@@ -132,9 +140,11 @@ class ChatRoute(Route):
return Response().error("Attachment not found").__dict__
file_path = attachment.path
real_file_path = os.path.realpath(file_path)
resolved_file_path = _resolve_path(file_path)
return await send_file(real_file_path, mimetype=attachment.mime_type)
return await send_file(
str(resolved_file_path), mimetype=attachment.mime_type
)
except (FileNotFoundError, OSError):
return Response().error("File access error").__dict__
@@ -715,10 +725,10 @@ class ChatRoute(Route):
try:
attachments = await self.db.get_attachments(attachment_ids)
for attachment in attachments:
if not os.path.exists(attachment.path):
if not await anyio.Path(attachment.path).exists():
continue
try:
os.remove(attachment.path)
await anyio.Path(attachment.path).unlink()
except OSError as e:
logger.warning(
f"Failed to delete attachment file {attachment.path}: {e}"
+15 -10
View File
@@ -6,6 +6,7 @@ import traceback
from pathlib import Path
from typing import Any
import anyio
from quart import request
from astrbot.core import astrbot_config, file_token_service, logger
@@ -40,6 +41,10 @@ from .util import (
MAX_FILE_BYTES = 500 * 1024 * 1024
def _resolve_path(path: Path) -> Path:
return path.resolve(strict=False)
def try_cast(value: Any, type_: str):
if type_ == "int":
try:
@@ -1105,8 +1110,8 @@ class ConfigRoute(Route):
if not files:
return Response().error("No files uploaded").__dict__
storage_root_path = Path(get_astrbot_plugin_data_path()).resolve(strict=False)
plugin_root_path = (storage_root_path / name).resolve(strict=False)
storage_root_path = _resolve_path(Path(get_astrbot_plugin_data_path()))
plugin_root_path = _resolve_path(storage_root_path / name)
try:
plugin_root_path.relative_to(storage_root_path)
except ValueError:
@@ -1133,7 +1138,7 @@ class ConfigRoute(Route):
continue
rel_path = f"files/{folder}/{filename}"
save_path = (plugin_root_path / rel_path).resolve(strict=False)
save_path = _resolve_path(plugin_root_path / rel_path)
try:
save_path.relative_to(plugin_root_path)
except ValueError:
@@ -1181,13 +1186,13 @@ class ConfigRoute(Route):
if not md:
return Response().error(f"Plugin {name} not found").__dict__
storage_root_path = Path(get_astrbot_plugin_data_path()).resolve(strict=False)
plugin_root_path = (storage_root_path / name).resolve(strict=False)
storage_root_path = _resolve_path(Path(get_astrbot_plugin_data_path()))
plugin_root_path = _resolve_path(storage_root_path / name)
try:
plugin_root_path.relative_to(storage_root_path)
except ValueError:
return Response().error("Invalid name parameter").__dict__
target_path = (plugin_root_path / rel_path).resolve(strict=False)
target_path = _resolve_path(plugin_root_path / rel_path)
try:
target_path.relative_to(plugin_root_path)
except ValueError:
@@ -1209,15 +1214,15 @@ class ConfigRoute(Route):
if not meta or meta.get("type") != "file":
return Response().error("Config item not found or not file type").__dict__
storage_root_path = Path(get_astrbot_plugin_data_path()).resolve(strict=False)
plugin_root_path = (storage_root_path / name).resolve(strict=False)
storage_root_path = _resolve_path(Path(get_astrbot_plugin_data_path()))
plugin_root_path = _resolve_path(storage_root_path / name)
try:
plugin_root_path.relative_to(storage_root_path)
except ValueError:
return Response().error("Invalid name parameter").__dict__
folder = config_key_to_folder(key_path)
target_dir = (plugin_root_path / "files" / folder).resolve(strict=False)
target_dir = _resolve_path(plugin_root_path / "files" / folder)
try:
target_dir.relative_to(plugin_root_path)
except ValueError:
@@ -1377,7 +1382,7 @@ class ConfigRoute(Route):
logo_file_path = os.path.join(plugin_dir, platform.logo_path)
# 检查文件是否存在并注册令牌
if os.path.exists(logo_file_path):
if await anyio.Path(logo_file_path).exists():
logo_token = await file_token_service.register_file(
logo_file_path,
expire_seconds=3600,
+15 -18
View File
@@ -7,6 +7,7 @@ from functools import cmp_to_key
from pathlib import Path
import aiohttp
import anyio
import psutil
from quart import request
@@ -22,6 +23,10 @@ from astrbot.core.utils.version_comparator import VersionComparator
from .route import Response, Route, RouteContext
def _resolve_path(path: str | Path) -> Path:
return Path(path).resolve(strict=False)
class StatRoute(Route):
def __init__(
self,
@@ -210,41 +215,33 @@ class StatRoute(Route):
filename = f"v{version}.md"
project_path = get_astrbot_path()
changelogs_dir = os.path.join(project_path, "changelogs")
changelog_path = os.path.join(changelogs_dir, filename)
# 规范化路径,防止符号链接攻击
changelog_path = os.path.realpath(changelog_path)
changelogs_dir = os.path.realpath(changelogs_dir)
changelogs_dir = _resolve_path(Path(project_path) / "changelogs")
changelog_path = _resolve_path(changelogs_dir / filename)
# 验证最终路径在预期的 changelogs 目录内(防止路径遍历)
# 确保规范化后的路径以 changelogs_dir 开头,且是目录内的文件
changelog_path_normalized = os.path.normpath(changelog_path)
changelogs_dir_normalized = os.path.normpath(changelogs_dir)
# 检查路径是否在预期目录内(必须是目录的子文件,不能是目录本身)
expected_prefix = changelogs_dir_normalized + os.sep
if not changelog_path_normalized.startswith(expected_prefix):
try:
changelog_path.relative_to(changelogs_dir)
except ValueError:
logger.warning(
f"Path traversal attempt detected: {version} -> {changelog_path}",
)
return Response().error("Invalid version format").__dict__
if not os.path.exists(changelog_path):
if not await anyio.Path(changelog_path).exists():
return (
Response()
.error(f"Changelog for version {version} not found")
.__dict__
)
if not os.path.isfile(changelog_path):
if not await anyio.Path(changelog_path).is_file():
return (
Response()
.error(f"Changelog for version {version} not found")
.__dict__
)
with open(changelog_path, encoding="utf-8") as f:
content = f.read()
async with await anyio.open_file(changelog_path, encoding="utf-8") as f:
content = await f.read()
return Response().ok({"content": content, "version": version}).__dict__
except Exception as e:
@@ -257,7 +254,7 @@ class StatRoute(Route):
project_path = get_astrbot_path()
changelogs_dir = os.path.join(project_path, "changelogs")
if not os.path.exists(changelogs_dir):
if not await anyio.Path(changelogs_dir).exists():
return Response().ok({"versions": []}).__dict__
versions = []
+4 -2
View File
@@ -5,6 +5,8 @@ import os
import sys
from pathlib import Path
import anyio
import runtime_bootstrap
from astrbot.core import LogBroker, LogManager, db_helper, logger
from astrbot.core.config.default import VERSION
@@ -69,13 +71,13 @@ async def check_dashboard_files(webui_dir: str | None = None):
"""下载管理面板文件"""
# 指定webui目录
if webui_dir:
if os.path.exists(webui_dir):
if await anyio.Path(webui_dir).exists():
logger.info(f"使用指定的 WebUI 目录: {webui_dir}")
return webui_dir
logger.warning(f"指定的 WebUI 目录 {webui_dir} 不存在,将使用默认逻辑。")
data_dist_path = os.path.join(get_astrbot_data_path(), "dist")
if os.path.exists(data_dist_path):
if await anyio.Path(data_dist_path).exists():
v = await get_dashboard_version()
if v is not None:
# 存在文件