P0-C: 安全修复

- 会话文件移出 web 根(data/ 目录),auth_secret 启动自动生成并持久化(0600)
- cookie 签名改 HMAC-SHA256 + compare_digest;会话表加锁 + 原子写
- serve_path 增加敏感文件黑名单(纵深防御)
- PyPI: 移除 ?url= 任意 URL 回退(SSRF),限制 scheme,缓存键防路径穿越
- v2 文件元数据/版本/缩略图端点全部加 is_safe_path;缩略图尺寸与像素上限
- api-docs/generate 限写 docs 目录;user/password 强制旧密码;GET /api/v2/config 脱敏
- 目录列表/错误页 HTML 转义(防存储型 XSS);Content-Disposition 文件名清洗
- 信号处理改优雅退出(移除 os._exit);启动时默认凭据安全警告
This commit is contained in:
HYC Fixer
2026-08-30 12:20:17 +08:00
parent 63d5addc25
commit 2a899c411e
6 changed files with 6938 additions and 6766 deletions
+2 -12
View File
@@ -90,18 +90,8 @@ class AdminAPI:
def _handle_sessions(self, handler, method, parts): def _handle_sessions(self, handler, method, parts):
"""处理会话管理""" """处理会话管理"""
if method == 'GET': if method == 'GET':
# 列出活跃会话 # 列出活跃会话(auth_manager 内部加锁遍历)
sessions = [] sessions = self.auth_manager.list_sessions()
for session_id, session in self.auth_manager.sessions.items():
if time.time() < session.expires_at:
sessions.append({
"session_id": session.session_id,
"user_id": session.user_id,
"level": session.level,
"created_at": session.created_at,
"expires_at": session.expires_at,
"last_activity": session.last_activity
})
handler.send_json_response({ handler.send_json_response({
"sessions": sessions, "sessions": sessions,
+68 -10
View File
@@ -11,7 +11,7 @@ from datetime import datetime
from .v1 import APIv1 from .v1 import APIv1
from .admin import AdminAPI from .admin import AdminAPI
from core.utils import format_file_size from core.utils import format_file_size, is_safe_path
from core.api_auth import require_auth, check_endpoint_auth from core.api_auth import require_auth, check_endpoint_auth
@@ -1621,6 +1621,9 @@ class APIv2(APIv1):
def api_get_file_metadata(self, handler, filename): def api_get_file_metadata(self, handler, filename):
"""获取文件元数据""" """获取文件元数据"""
full_path = os.path.join(self.config['base_dir'], filename) full_path = os.path.join(self.config['base_dir'], filename)
if not is_safe_path(self.config['base_dir'], full_path):
handler.send_json_response({"error": "Access denied"}, 403)
return
if not os.path.exists(full_path): if not os.path.exists(full_path):
handler.send_json_response({"error": "File not found"}, 404) handler.send_json_response({"error": "File not found"}, 404)
return return
@@ -1657,6 +1660,9 @@ class APIv2(APIv1):
try: try:
metadata = json.loads(handler.rfile.read(content_length)) metadata = json.loads(handler.rfile.read(content_length))
full_path = os.path.join(self.config['base_dir'], filename) full_path = os.path.join(self.config['base_dir'], filename)
if not is_safe_path(self.config['base_dir'], full_path):
handler.send_json_response({"error": "Access denied"}, 403)
return
metadata_file = full_path + '.meta' metadata_file = full_path + '.meta'
with open(metadata_file, 'w', encoding='utf-8') as f: with open(metadata_file, 'w', encoding='utf-8') as f:
@@ -1676,6 +1682,8 @@ class APIv2(APIv1):
results = {} results = {}
for path in paths: for path in paths:
full_path = os.path.join(self.config['base_dir'], path) full_path = os.path.join(self.config['base_dir'], path)
if not is_safe_path(self.config['base_dir'], full_path):
continue
if os.path.exists(full_path): if os.path.exists(full_path):
metadata_file = full_path + '.meta' metadata_file = full_path + '.meta'
metadata = {} metadata = {}
@@ -1702,6 +1710,9 @@ class APIv2(APIv1):
for path, metadata in data.items(): for path, metadata in data.items():
full_path = os.path.join(self.config['base_dir'], path) full_path = os.path.join(self.config['base_dir'], path)
if not is_safe_path(self.config['base_dir'], full_path):
results[path] = {"success": False, "error": "Access denied"}
continue
metadata_file = full_path + '.meta' metadata_file = full_path + '.meta'
try: try:
with open(metadata_file, 'w', encoding='utf-8') as f: with open(metadata_file, 'w', encoding='utf-8') as f:
@@ -1721,6 +1732,9 @@ class APIv2(APIv1):
# 这里可以实现版本控制系统 # 这里可以实现版本控制系统
# 简化实现:检查备份文件 # 简化实现:检查备份文件
full_path = os.path.join(self.config['base_dir'], filename) full_path = os.path.join(self.config['base_dir'], filename)
if not is_safe_path(self.config['base_dir'], full_path):
handler.send_json_response({"error": "Access denied"}, 403)
return
versions = [] versions = []
# 查找备份文件 # 查找备份文件
@@ -1750,6 +1764,9 @@ class APIv2(APIv1):
"""创建文件版本(备份)""" """创建文件版本(备份)"""
import time import time
full_path = os.path.join(self.config['base_dir'], filename) full_path = os.path.join(self.config['base_dir'], filename)
if not is_safe_path(self.config['base_dir'], full_path):
handler.send_json_response({"error": "Access denied"}, 403)
return
if not os.path.exists(full_path): if not os.path.exists(full_path):
handler.send_json_response({"error": "File not found"}, 404) handler.send_json_response({"error": "File not found"}, 404)
return return
@@ -1775,6 +1792,9 @@ class APIv2(APIv1):
import mimetypes import mimetypes
full_path = os.path.join(self.config['base_dir'], filename) full_path = os.path.join(self.config['base_dir'], filename)
if not is_safe_path(self.config['base_dir'], full_path):
handler.send_json_response({"error": "Access denied"}, 403)
return
if not os.path.exists(full_path): if not os.path.exists(full_path):
handler.send_json_response({"error": "File not found"}, 404) handler.send_json_response({"error": "File not found"}, 404)
return return
@@ -1784,8 +1804,17 @@ class APIv2(APIv1):
handler.send_json_response({"error": "Not an image file"}, 400) handler.send_json_response({"error": "Not an image file"}, 400)
return return
width = int(query_params.get('width', ['200'])[0]) try:
height = int(query_params.get('height', ['200'])[0]) width = min(max(int(query_params.get('width', ['200'])[0]), 16), 2048)
height = min(max(int(query_params.get('height', ['200'])[0]), 16), 2048)
except ValueError:
width, height = 200, 200
# 限制解码像素上限,防解压炸弹
try:
from PIL import Image as _PILImage
_PILImage.MAX_IMAGE_PIXELS = 50_000_000
except Exception:
pass
try: try:
from PIL import Image from PIL import Image
@@ -3315,17 +3344,16 @@ class APIv2(APIv1):
old_password = data.get('old_password') old_password = data.get('old_password')
new_password = data.get('new_password') new_password = data.get('new_password')
if not username or not new_password: if not username or not new_password or not old_password:
handler.send_json_response({"error": "缺少必要参数"}, 400) handler.send_json_response({"error": "缺少必要参数(username / old_password / new_password)"}, 400)
return return
# 获取配置中的账号密码 # 获取配置中的账号密码
config_user = config.get('auth_user', '') config_user = config.get('auth_user', '')
config_pass = config.get('auth_pass', '') config_pass = config.get('auth_pass', '')
# 验证旧密码(优先验证数据库,没有则验证配置文件) # 验证旧密码(优先验证数据库,没有则验证配置文件)—— 强制要求,防止无旧密码改密
if old_password: user = db.get_user(username) if db else None
user = db.get_user(username)
if user: if user:
# 验证数据库密码 # 验证数据库密码
if not db.verify_password(old_password, user['password_hash']): if not db.verify_password(old_password, user['password_hash']):
@@ -3417,12 +3445,28 @@ class APIv2(APIv1):
}, 404) }, 404)
return return
# 脱敏: 不返回 settings.json 原文(含 auth_pass/auth_token 等密钥),
# 只返回精选非敏感字段
try:
with open(config_path, 'r', encoding='utf-8') as f: with open(config_path, 'r', encoding='utf-8') as f:
config_content = f.read() raw = json.load(f)
except Exception:
raw = {}
safe_keys = [
'server_name', 'host', 'port', 'base_dir', 'api_version',
'directory_listing', 'enable_stats', 'show_hash', 'ignore_hidden',
'enable_range', 'max_workers', 'timeout', 'max_upload_size',
'enable_ws', 'enable_sse', 'enable_monitor', 'monitor_interval',
'enable_sync', 'enable_mirrors', 'auth_type', 'log_level',
'cache_size', 'cache_ttl', 'sort_by', 'sort_reverse',
'max_search_results', 'session_timeout', 'sync_interval',
]
safe_config = {k: raw.get(k) for k in safe_keys if k in raw}
handler.send_json_response({ handler.send_json_response({
"success": True, "success": True,
"config": config_content, "config": safe_config,
"path": config_path, "path": config_path,
"filename": os.path.basename(config_path) "filename": os.path.basename(config_path)
}) })
@@ -4397,6 +4441,20 @@ class APIv2(APIv1):
filepath = 'docs/api-docs.json' filepath = 'docs/api-docs.json'
format = 'json' format = 'json'
# 路径安全校验: 只允许写入项目 docs/ 或 api/docs/ 目录
project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
allowed_dirs = [
os.path.realpath(os.path.join(project_root, 'docs')),
os.path.realpath(os.path.join(project_root, 'api', 'docs')),
]
resolved = os.path.realpath(os.path.join(project_root, filepath))
if not any(resolved == d or resolved.startswith(d + os.sep) for d in allowed_dirs):
handler.send_json_response({
'success': False,
'error': 'Invalid filepath: must be under docs/'
}, 403)
return
# 生成并保存文档 # 生成并保存文档
saved_path = save_api_docs(self.config, filepath, format) saved_path = save_api_docs(self.config, filepath, format)
+72 -10
View File
@@ -38,6 +38,7 @@ class APIAuthManager:
def __init__(self, config: dict = None): def __init__(self, config: dict = None):
self.config = config or {} self.config = config or {}
self.sessions: Dict[str, AuthSession] = {} self.sessions: Dict[str, AuthSession] = {}
self._lock = threading.Lock()
# 确定基础目录(用于保存会话文件) # 确定基础目录(用于保存会话文件)
base_dir = config.get('base_dir', '.') if config else '.' base_dir = config.get('base_dir', '.') if config else '.'
@@ -48,9 +49,20 @@ class APIAuthManager:
elif base_dir == '.': elif base_dir == '.':
base_dir = os.getcwd() base_dir = os.getcwd()
# 数据目录:默认放到 base_dir 同级 data/ 下,
# 避免 auth_sessions.json 落入 web 静态根目录被公开下载
data_dir = config.get('data_dir') if config else None
if not data_dir:
data_dir = os.path.join(os.path.dirname(os.path.abspath(base_dir)), 'data')
try:
os.makedirs(data_dir, exist_ok=True)
except OSError:
data_dir = base_dir
self.data_dir = data_dir
# 会话文件路径 # 会话文件路径
sessions_filename = config.get('auth_sessions_file', 'auth_sessions.json') if config else 'auth_sessions.json' sessions_filename = config.get('auth_sessions_file', 'auth_sessions.json') if config else 'auth_sessions.json'
self.sessions_file = os.path.join(base_dir, sessions_filename) self.sessions_file = os.path.join(data_dir, sessions_filename)
# 会话超时时间 # 会话超时时间
self.session_timeout = config.get('auth_session_timeout', 3600) if config else 3600 self.session_timeout = config.get('auth_session_timeout', 3600) if config else 3600
@@ -63,9 +75,33 @@ class APIAuthManager:
self.ip_whitelist = config.get('ip_whitelist', []) if config else [] self.ip_whitelist = config.get('ip_whitelist', []) if config else []
self.ip_whitelist_enabled = config.get('ip_whitelist_enabled', False) if config else False self.ip_whitelist_enabled = config.get('ip_whitelist_enabled', False) if config else False
# Cookie 签名密钥:优先用配置;缺失则从文件读取或生成并持久化
self.auth_secret = self._load_or_create_secret(config)
# 加载已保存的会话 # 加载已保存的会话
self._load_sessions() self._load_sessions()
def _load_or_create_secret(self, config) -> str:
"""获取或生成 auth_secret(持久化到数据目录,避免默认密钥公开可伪造)"""
secret = config.get('auth_secret') if config else None
if secret:
return secret
secret_file = os.path.join(self.data_dir, 'auth_secret.key')
try:
if os.path.exists(secret_file):
with open(secret_file, 'r', encoding='utf-8') as f:
secret = f.read().strip()
if secret:
return secret
secret = secrets.token_hex(32)
with open(secret_file, 'w', encoding='utf-8') as f:
f.write(secret)
os.chmod(secret_file, 0o600)
except Exception as e:
print(f"警告: 持久化 auth_secret 失败: {e}")
return secret
@property @property
def db(self): def db(self):
"""动态获取数据库实例""" """动态获取数据库实例"""
@@ -182,6 +218,7 @@ class APIAuthManager:
def validate_cookie(self, cookie_value: str, client_ip: str = None) -> Optional[dict]: def validate_cookie(self, cookie_value: str, client_ip: str = None) -> Optional[dict]:
"""验证认证Cookie""" """验证认证Cookie"""
import hmac
if not cookie_value: if not cookie_value:
return None return None
@@ -191,6 +228,7 @@ class APIAuthManager:
session_id, timestamp, signature = parts session_id, timestamp, signature = parts
with self._lock:
session = self.sessions.get(session_id) session = self.sessions.get(session_id)
if not session: if not session:
return None return None
@@ -200,7 +238,7 @@ class APIAuthManager:
return None return None
expected_sig = self._generate_signature(session_id, timestamp, session.user_id) expected_sig = self._generate_signature(session_id, timestamp, session.user_id)
if signature != expected_sig: if not hmac.compare_digest(signature, expected_sig):
return None return None
session.last_activity = time.time() session.last_activity = time.time()
@@ -218,6 +256,7 @@ class APIAuthManager:
if not session_id: if not session_id:
return None return None
with self._lock:
session = self.sessions.get(session_id) session = self.sessions.get(session_id)
if not session: if not session:
return None return None
@@ -355,6 +394,7 @@ class APIAuthManager:
permissions=permissions or ['*'] permissions=permissions or ['*']
) )
with self._lock:
self.sessions[session_id] = session self.sessions[session_id] = session
self._save_sessions() self._save_sessions()
@@ -371,19 +411,37 @@ class APIAuthManager:
def destroy_session(self, session_id: str) -> bool: def destroy_session(self, session_id: str) -> bool:
"""销毁会话""" """销毁会话"""
with self._lock:
if session_id in self.sessions: if session_id in self.sessions:
del self.sessions[session_id] del self.sessions[session_id]
self._save_sessions() self._save_sessions()
return True return True
return False return False
def list_sessions(self) -> List[dict]:
"""列出活跃会话(加锁遍历,供管理接口使用)"""
now = time.time()
with self._lock:
sessions = []
for session_id, session in self.sessions.items():
if now < session.expires_at:
sessions.append({
"session_id": session.session_id,
"user_id": session.user_id,
"level": session.level,
"created_at": session.created_at,
"expires_at": session.expires_at,
"last_activity": session.last_activity
})
return sessions
# === 内部方法 === # === 内部方法 ===
def _generate_signature(self, session_id: str, timestamp: str, user_id: str) -> str: def _generate_signature(self, session_id: str, timestamp: str, user_id: str) -> str:
"""生成签名""" """生成签名(HMAC-SHA256,使用持久化的 auth_secret)"""
secret = self.config.get('auth_secret', 'default_secret_change_me') import hmac
data = f"{session_id}.{timestamp}.{user_id}.{secret}" data = f"{session_id}.{timestamp}.{user_id}".encode('utf-8')
return hashlib.sha256(data.encode()).hexdigest()[:32] return hmac.new(self.auth_secret.encode('utf-8'), data, hashlib.sha256).hexdigest()[:32]
def _load_sessions(self): def _load_sessions(self):
"""加载会话""" """加载会话"""
@@ -409,7 +467,8 @@ class APIAuthManager:
print(f"加载会话失败: {e}") print(f"加载会话失败: {e}")
def _save_sessions(self): def _save_sessions(self):
"""保存会话""" """保存会话(原子写: 临时文件 + os.replace)"""
with self._lock:
data = [] data = []
for session in self.sessions.values(): for session in self.sessions.values():
data.append({ data.append({
@@ -423,13 +482,16 @@ class APIAuthManager:
}) })
try: try:
with open(self.sessions_file, 'w', encoding='utf-8') as f: tmp_file = self.sessions_file + '.tmp'
with open(tmp_file, 'w', encoding='utf-8') as f:
json.dump(data, f, ensure_ascii=False, indent=2) json.dump(data, f, ensure_ascii=False, indent=2)
except Exception: os.replace(tmp_file, self.sessions_file)
pass except Exception as e:
print(f"警告: 保存会话失败: {e}")
def get_stats(self) -> dict: def get_stats(self) -> dict:
"""获取认证统计""" """获取认证统计"""
with self._lock:
active_sessions = sum( active_sessions = sum(
1 for s in self.sessions.values() 1 for s in self.sessions.values()
if time.time() < s.expires_at if time.time() < s.expires_at
+38 -14
View File
@@ -780,6 +780,16 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
self.send_error(500) self.send_error(500)
return return
# 敏感文件黑名单(纵深防御:禁止下载会话/密钥等文件)
SENSITIVE_FILES = {
'auth_sessions.json', 'auth_secret.key', 'auth_token.txt',
'settings.json', 'sync_state.json', 'webhooks.json',
}
base_name = os.path.basename(rel_path).lower()
if base_name in SENSITIVE_FILES:
self.send_error(403, "Access denied")
return
file_path = os.path.join(self.config['base_dir'], rel_path) file_path = os.path.join(self.config['base_dir'], rel_path)
if not is_safe_path(self.config['base_dir'], file_path): if not is_safe_path(self.config['base_dir'], file_path):
@@ -888,11 +898,13 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
# 动态计算列数 # 动态计算列数
colspan = 4 if self.config.get('show_hash') else 3 colspan = 4 if self.config.get('show_hash') else 3
import html as _html
title_safe = _html.escape(str(title))
html = f"""<!DOCTYPE html> html = f"""<!DOCTYPE html>
<html lang="zh-CN"> <html lang="zh-CN">
<head> <head>
<meta charset="utf-8"> <meta charset="utf-8">
<title>{title}</title> <title>{title_safe}</title>
<meta name="viewport" content="width=device-width, initial-scale=1.0"> <meta name="viewport" content="width=device-width, initial-scale=1.0">
<style> <style>
body {{ body {{
@@ -986,7 +998,7 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
<div class="container"> <div class="container">
<h1>{title}</h1> <h1>{title}</h1>
<div class="breadcrumb"> <div class="breadcrumb">
{' / '.join(f'<a href="{crumb["path"]}">{crumb["name"]}</a>' for crumb in breadcrumbs)} {' / '.join(f'<a href="{_html.escape(crumb["path"], quote=True)}">{_html.escape(crumb["name"])}</a>' for crumb in breadcrumbs)}
</div> </div>
<table> <table>
<thead> <thead>
@@ -1001,28 +1013,38 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
# 修复:只在非根目录显示上一级目录链接 # 修复:只在非根目录显示上一级目录链接
if rel_dir: # 如果不是根目录 if rel_dir: # 如果不是根目录
html += f'<tr class="dir"><td colspan="{colspan}"><a href="{parent_path}">../</a></td></tr>\n' html += f'<tr class="dir"><td colspan="{colspan}"><a href="{_html.escape(parent_path, quote=True)}">../</a></td></tr>\n'
for item in items: for item in items:
item_name_safe = _html.escape(str(item["name"]))
item_path_safe = _html.escape(item["path"], quote=True)
html += f'<tr class="{"dir" if item["is_dir"] else "file"}">' html += f'<tr class="{"dir" if item["is_dir"] else "file"}">'
html += f'<td><a href="/{item["path"]}">{item["name"]}{" /" if item["is_dir"] else ""}</a></td>' html += f'<td><a href="/{item_path_safe}">{item_name_safe}{" /" if item["is_dir"] else ""}</a></td>'
html += f'<td class="modified">{item["modified"]}</td>' html += f'<td class="modified">{_html.escape(str(item["modified"]))}</td>'
html += f'<td class="size">{item["size"]}</td>' html += f'<td class="size">{_html.escape(str(item["size"]))}</td>'
if self.config.get('show_hash'): if self.config.get('show_hash'):
html += f'<td class="sha256">{item["sha256"]}</td>' html += f'<td class="sha256">{_html.escape(str(item["sha256"]))}</td>'
html += '</tr>\n' html += '</tr>\n'
html += f""" html += f"""
</tbody> </tbody>
</table> </table>
<div class="server-info"> <div class="server-info">
<p>Files: {len(items)} | {self.config.get("server_name", "Mirror Server")}</p> <p>Files: {len(items)} | {_html.escape(str(self.config.get("server_name", "Mirror Server")))}</p>
</div> </div>
</div> </div>
</body> </body>
</html>""" </html>"""
return html return html
@staticmethod
def _safe_disposition_filename(file_path):
"""清洗 Content-Disposition 文件名,防响应头注入(引号/CRLF/控制字符)"""
name = os.path.basename(file_path)
# 去掉引号与换行等危险字符
name = re.sub(r'["\r\n\x00-\x1f]', '_', name)
return name
def send_file_headers(self, file_path): def send_file_headers(self, file_path):
"""发送文件头信息(用于HEAD请求)""" """发送文件头信息(用于HEAD请求)"""
if not os.path.exists(file_path) or not os.path.isfile(file_path): if not os.path.exists(file_path) or not os.path.isfile(file_path):
@@ -1037,7 +1059,7 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
self.send_header("Content-Type", mime_type) self.send_header("Content-Type", mime_type)
self.send_header("Content-Length", str(file_size)) self.send_header("Content-Length", str(file_size))
self.send_header("Content-Disposition", self.send_header("Content-Disposition",
f'attachment; filename="{os.path.basename(file_path)}"') f'attachment; filename="{self._safe_disposition_filename(file_path)}"')
self.send_header("Accept-Ranges", "bytes") self.send_header("Accept-Ranges", "bytes")
self.send_header("Cache-Control", "public, max-age=3600") self.send_header("Cache-Control", "public, max-age=3600")
self.send_header( self.send_header(
@@ -1092,10 +1114,10 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
self.send_header("Content-Length", str(content_length)) self.send_header("Content-Length", str(content_length))
# HTML 文件直接在浏览器中显示,不强制下载 # HTML 文件直接在浏览器中显示,不强制下载
if mime_type == 'text/html': if mime_type == 'text/html':
self.send_header("Content-Disposition", f'inline; filename="{os.path.basename(file_path)}"') self.send_header("Content-Disposition", f'inline; filename="{self._safe_disposition_filename(file_path)}"')
else: else:
self.send_header("Content-Disposition", self.send_header("Content-Disposition",
f'attachment; filename="{os.path.basename(file_path)}"') f'attachment; filename="{self._safe_disposition_filename(file_path)}"')
self.send_header("Accept-Ranges", "bytes") self.send_header("Accept-Ranges", "bytes")
self.send_header("Cache-Control", "public, max-age=3600") self.send_header("Cache-Control", "public, max-age=3600")
self.send_header("Last-Modified", self.date_time_string(os.path.getmtime(file_path))) self.send_header("Last-Modified", self.date_time_string(os.path.getmtime(file_path)))
@@ -1180,7 +1202,7 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
self.send_header("Content-Type", mime_type) self.send_header("Content-Type", mime_type)
self.send_header("Content-Length", str(file_size)) self.send_header("Content-Length", str(file_size))
self.send_header("Content-Disposition", self.send_header("Content-Disposition",
f'attachment; filename="{os.path.basename(file_path)}"') f'attachment; filename="{self._safe_disposition_filename(file_path)}"')
self.send_header("Accept-Ranges", "bytes") self.send_header("Accept-Ranges", "bytes")
self.send_header("Transfer-Encoding", "chunked") self.send_header("Transfer-Encoding", "chunked")
self.end_headers() self.end_headers()
@@ -1231,13 +1253,15 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
if message is None: if message is None:
message = error_messages.get(code, "未知错误") message = error_messages.get(code, "未知错误")
import html as _html
message_safe = _html.escape(str(message))
error_page = f""" error_page = f"""
<!DOCTYPE html> <!DOCTYPE html>
<html lang="zh-CN"> <html lang="zh-CN">
<head> <head>
<meta charset="UTF-8"> <meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0"> <meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>{code} {message}</title> <title>{code} {message_safe}</title>
<style> <style>
body {{ body {{
font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif; font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif;
@@ -1288,7 +1312,7 @@ class MirrorServerHandler(BaseHTTPRequestHandler):
<body> <body>
<div class="error-container"> <div class="error-container">
<h1 class="error-code">{code}</h1> <h1 class="error-code">{code}</h1>
<h2 class="error-message">{message}</h2> <h2 class="error-message">{message_safe}</h2>
<p class="error-description">请求的页面遇到问题,请稍后重试.</p> <p class="error-description">请求的页面遇到问题,请稍后重试.</p>
<a href="/" class="home-link">返回首页</a> <a href="/" class="home-link">返回首页</a>
</div> </div>
+22 -4
View File
@@ -36,10 +36,9 @@ from core.optimization import (
def signal_handler(signum, _frame): def signal_handler(signum, _frame):
"""处理退出信号""" """处理退出信号(优雅退出:触发 KeyboardInterrupt 走正常关闭流程)"""
print(f"\n收到信号 {signum},正在关闭服务器...") print(f"\n收到信号 {signum},正在优雅关闭服务器...")
import os raise KeyboardInterrupt
os._exit(0)
def parse_arguments(): def parse_arguments():
@@ -587,6 +586,21 @@ def main():
# 设置服务器启动时间(用于计算运行时间) # 设置服务器启动时间(用于计算运行时间)
config['start_time'] = time.time() config['start_time'] = time.time()
# 默认凭据警告
try:
if config.get('auth_type') == 'basic' and config.get('auth_pass') in (None, '', 'admin123'):
print("\n" + "!" * 60)
print("! 安全警告: 正在使用默认/空认证密码 (admin/admin123)")
print("! 请立即修改 settings.json 中的 auth_pass 或使用 --auth-pass 指定")
print("!" * 60)
if config.get('auth_type') == 'token' and config.get('auth_token') in (None, '', 'admin_token_123456'):
print("\n" + "!" * 60)
print("! 安全警告: 正在使用默认/空认证 token")
print("! 请立即修改 settings.json 中的 auth_token 或使用 --auth-token 指定")
print("!" * 60)
except Exception:
pass
# 创建并启动服务器 # 创建并启动服务器
try: try:
server = MirrorServer(config) server = MirrorServer(config)
@@ -595,6 +609,10 @@ def main():
else: else:
print("服务器启动失败") print("服务器启动失败")
sys.exit(1) sys.exit(1)
except KeyboardInterrupt:
# 信号处理已触发优雅关闭(server.stop + atexit cleanup 已执行)
print("\n服务器已正常退出")
sys.exit(0)
except Exception as e: except Exception as e:
print(f"错误: {e}") print(f"错误: {e}")
import traceback import traceback
+29 -9
View File
@@ -413,16 +413,16 @@ class PyPIMirror:
# 尝试官方源 # 尝试官方源
possible_urls.append(f"https://files.pythonhosted.org/packages/{actual_filename}") possible_urls.append(f"https://files.pythonhosted.org/packages/{actual_filename}")
# 如果有查询参数中的URL,也尝试 # 注意: 不再支持 ?url= 参数指定任意上游 URL(SSRF 风险),
if 'url' in query_params: # 只允许从配置的上游与官方源获取
possible_urls.insert(0, urllib.parse.unquote(query_params['url'][0]))
data = None data = None
last_error = None last_error = None
for url in possible_urls: for url in possible_urls:
if not url.startswith(('http://', 'https://')):
continue
try: try:
import sys
req = urllib.request.Request(url) req = urllib.request.Request(url)
req.add_header('User-Agent', 'PyPI-Mirror/1.0') req.add_header('User-Agent', 'PyPI-Mirror/1.0')
@@ -703,6 +703,8 @@ class PyPIMirror:
return None return None
cache_path = self._get_cache_path(cache_key) cache_path = self._get_cache_path(cache_key)
if cache_path is None:
return None
meta_path = cache_path + '.meta' meta_path = cache_path + '.meta'
if not os.path.exists(cache_path): if not os.path.exists(cache_path):
@@ -726,6 +728,8 @@ class PyPIMirror:
def _set_cache(self, cache_key: str, data: bytes): def _set_cache(self, cache_key: str, data: bytes):
"""设置缓存""" """设置缓存"""
cache_path = self._get_cache_path(cache_key) cache_path = self._get_cache_path(cache_key)
if cache_path is None:
return
meta_path = cache_path + '.meta' meta_path = cache_path + '.meta'
os.makedirs(os.path.dirname(cache_path), exist_ok=True) os.makedirs(os.path.dirname(cache_path), exist_ok=True)
@@ -746,13 +750,29 @@ class PyPIMirror:
except Exception as e: except Exception as e:
print(f"PyPI缓存写入失败: {e}") print(f"PyPI缓存写入失败: {e}")
def _get_cache_path(self, cache_key: str) -> str: def _sanitize_cache_key(self, cache_key: str):
"""获取缓存路径""" """清洗缓存键,拒绝路径穿越(返回 None 表示不安全)"""
safe_key = cache_key.replace('\\', '/')
parts = []
for seg in safe_key.split('/'):
if seg in ('', '.'):
continue
if seg == '..':
return None # 路径穿越
parts.append(seg)
return '/'.join(parts)
def _get_cache_path(self, cache_key: str):
"""获取缓存路径(防路径穿越,不安全返回 None)"""
# cache_key 格式: packages/fe/df/88ccbee.../filename # cache_key 格式: packages/fe/df/88ccbee.../filename
# 存储路径: downloads/pypi-cn/packages/fe/df/88ccbee.../filename # 存储路径: downloads/pypi-cn/packages/fe/df/88ccbee.../filename
# 确保使用正斜杠 safe_key = self._sanitize_cache_key(cache_key)
safe_key = cache_key.replace('\\', '/') if safe_key is None:
return os.path.join(self.storage_dir, safe_key) return None
path = os.path.join(self.storage_dir, safe_key)
if not os.path.realpath(path).startswith(os.path.realpath(self.storage_dir) + os.sep):
return None
return path
def get_cache_stats(self) -> dict: def get_cache_stats(self) -> dict:
"""获取缓存统计""" """获取缓存统计"""