#!/usr/bin/env python3 # -*- coding: utf-8 -*- """ API认证模块 使用数据库进行认证,支持: - none: 无认证 - basic: Basic Auth(用户名密码) - token: Token 认证(登录生成的token) """ import os import sys import json import hashlib import time import secrets from typing import Dict, List, Optional from dataclasses import dataclass from functools import wraps @dataclass class AuthSession: """认证会话""" session_id: str user_id: str level: str created_at: float expires_at: float last_activity: float permissions: List[str] class APIAuthManager: """API认证管理器""" def __init__(self, config: dict = None): self.config = config or {} self.sessions: Dict[str, AuthSession] = {} # 确定基础目录(用于保存会话文件) base_dir = config.get('base_dir', '.') if config else '.' # 检测是否是 PyInstaller 打包环境 if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'): base_dir = os.path.dirname(os.path.abspath(sys.executable)) elif base_dir == '.': base_dir = os.getcwd() # 会话文件路径 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.session_timeout = config.get('auth_session_timeout', 3600) if config else 3600 # Cookie名称 self.cookie_name = 'hyc_auth' self.cookie_max_age = config.get('auth_cookie_max_age', 86400) if config else 86400 # IP 白名单 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._load_sessions() @property def db(self): """动态获取数据库实例""" return self.config.get('_db_instance') def _get_client_ip(self, handler) -> str: """从 handler 获取客户端 IP""" try: forwarded = handler.headers.get('X-Forwarded-For') if forwarded: return forwarded.split(',')[0].strip() return handler.client_address[0] except Exception: return None def _check_ip_whitelist(self, ip: str) -> bool: """检查 IP 是否在白名单中""" if not self.ip_whitelist_enabled or not self.ip_whitelist: return True if not ip: return False import ipaddress for pattern in self.ip_whitelist: try: if '/' in pattern: network = ipaddress.ip_network(pattern, strict=False) if ipaddress.ip_address(ip) in network: return True elif pattern == ip: return True except ValueError: continue return False # === Token 验证 === def validate_token(self, token: str, client_ip: str = None) -> Optional[dict]: """验证 token(从数据库)""" if not token: return None # 优先从数据库验证 if self.db: user = self.db.get_user_by_token(token) if user: # 检查 token 是否过期 if user.get('token_expires_at') and time.time() > user['token_expires_at']: return {"valid": False, "reason": "Token已过期"} # 检查用户是否启用 if not user.get('enabled', True): return {"valid": False, "reason": "用户已被禁用"} return { "valid": True, "key_id": f"user_{user['id']}", "name": user['username'], "level": user.get('role', 'admin'), "permissions": ["*"], "user_id": user['id'], "username": user['username'] } return {"valid": False, "reason": "无效的Token"} # === Basic Auth 验证 === def validate_basic_auth(self, username: str, password: str, client_ip: str = None) -> dict: """验证 Basic Auth 用户名密码(从数据库)""" # IP 白名单检查 if not self._check_ip_whitelist(client_ip): if self.db: self.db.add_login_log(username, client_ip, 'failed', 'IP不在白名单') return {"valid": False, "reason": "IP不在白名单内"} auth_type = self.config.get('auth_type', 'none') # 如果认证类型为 none,任何用户都可以通过 if auth_type == 'none': return { "valid": True, "user_id": 0, "username": username or 'anonymous', "level": "admin", "key_id": "anonymous", "name": f"Anonymous - {username or 'anonymous'}", "permissions": ["*"] } # 从数据库验证 if self.db: result = self.db.verify_user(username, password) if result.get('valid'): if self.db: self.db.add_login_log(username, client_ip, 'success', '数据库验证') return { "valid": True, "user_id": result.get('user_id'), "username": username, "level": result.get('role', 'admin'), "key_id": f"user_{result.get('user_id')}", "name": f"User - {username}", "permissions": ["*"] } else: if self.db: self.db.add_login_log(username, client_ip, 'failed', result.get('reason', '验证失败')) return {"valid": False, "reason": result.get('reason', '用户名或密码错误')} return {"valid": False, "reason": "数据库不可用"} # === Cookie 验证 === def validate_cookie(self, cookie_value: str, client_ip: str = None) -> Optional[dict]: """验证认证Cookie""" if not cookie_value: return None parts = cookie_value.split('.') if len(parts) != 3: return None session_id, timestamp, signature = parts session = self.sessions.get(session_id) if not session: return None if time.time() > session.expires_at: del self.sessions[session_id] return None expected_sig = self._generate_signature(session_id, timestamp, session.user_id) if signature != expected_sig: return None session.last_activity = time.time() return { "valid": True, "session_id": session_id, "user_id": session.user_id, "level": session.level, "permissions": session.permissions } def validate_session_id(self, session_id: str) -> Optional[dict]: """验证会话ID""" if not session_id: return None session = self.sessions.get(session_id) if not session: return None if time.time() > session.expires_at: del self.sessions[session_id] return None session.last_activity = time.time() return { "valid": True, "session_id": session_id, "user_id": session.user_id, "level": session.level, "permissions": session.permissions } # === 请求验证 === def validate_request(self, handler, required_level: str = "admin") -> dict: """ 验证请求的认证状态 支持的认证方式: 1. Authorization: Bearer 2. Authorization: Basic 3. X-API-Key: 4. Cookie: hyc_auth= 5. ?key= """ import base64 auth_header = handler.headers.get('Authorization') api_key = handler.headers.get('X-API-Key') cookie = handler.headers.get('Cookie', '') client_ip = handler.client_address[0] if hasattr(handler, 'client_address') else None # 提取cookie值 cookie_value = None for c in cookie.split(';'): c = c.strip() if c.startswith(f'{self.cookie_name}='): cookie_value = c[len(self.cookie_name)+1:] break # 获取查询参数中的key parsed_path = handler.path.split('?') query_key = None if len(parsed_path) > 1: from urllib.parse import parse_qs query = parse_qs(parsed_path[1]) query_key = query.get('key', [None])[0] # 1. Bearer Token if auth_header and auth_header.startswith('Bearer '): token = auth_header[7:] result = self.validate_token(token, client_ip) if result and result.get('valid'): return {"authenticated": True, "method": "bearer", **result} # 2. Basic Auth if auth_header and auth_header.startswith('Basic '): try: credentials = base64.b64decode(auth_header[6:]).decode('utf-8') if ':' in credentials: username, password = credentials.split(':', 1) result = self.validate_basic_auth(username, password, client_ip) if result and result.get('valid'): return {"authenticated": True, "method": "basic", **result} except Exception: pass # 3. API Key Header if api_key: result = self.validate_token(api_key, client_ip) if result and result.get('valid'): return {"authenticated": True, "method": "api_key", **result} # 4. Cookie if cookie_value: result = self.validate_cookie(cookie_value, client_ip) if result and result.get('valid'): return {"authenticated": True, "method": "cookie", **result} # 5. Query Parameter if query_key: result = self.validate_token(query_key, client_ip) if result and result.get('valid'): return {"authenticated": True, "method": "query", **result} # 未认证 return { "authenticated": False, "error": "Authentication required", "required_level": required_level } def check_permission(self, auth_result: dict, permission: str) -> bool: """检查是否有权限访问特定API""" if not auth_result.get('authenticated'): return False permissions = auth_result.get('permissions', []) if '*' in permissions: return True if permission in permissions: return True for p in permissions: if p.endswith('*'): prefix = p.rstrip('*') if permission.startswith(prefix): return True return False # === 会话管理 === def create_session(self, user_id: str, level: str, permissions: List[str] = None) -> dict: """创建认证会话""" session_id = secrets.token_hex(32) timestamp = time.time() session = AuthSession( session_id=session_id, user_id=user_id, level=level, created_at=timestamp, expires_at=timestamp + self.session_timeout, last_activity=timestamp, permissions=permissions or ['*'] ) self.sessions[session_id] = session self._save_sessions() signature = self._generate_signature(session_id, timestamp, user_id) cookie_value = f"{session_id}.{timestamp}.{signature}" return { "session_id": session_id, "cookie_name": self.cookie_name, "cookie_value": cookie_value, "cookie_max_age": self.cookie_max_age, "expires": timestamp + self.session_timeout } def destroy_session(self, session_id: str) -> bool: """销毁会话""" if session_id in self.sessions: del self.sessions[session_id] self._save_sessions() return True return False # === 内部方法 === def _generate_signature(self, session_id: str, timestamp: str, user_id: str) -> str: """生成签名""" secret = self.config.get('auth_secret', 'default_secret_change_me') data = f"{session_id}.{timestamp}.{user_id}.{secret}" return hashlib.sha256(data.encode()).hexdigest()[:32] def _load_sessions(self): """加载会话""" if os.path.exists(self.sessions_file): try: with open(self.sessions_file, 'r', encoding='utf-8') as f: data = json.load(f) now = time.time() for item in data: if item.get('expires_at') and now > item['expires_at']: continue session = AuthSession( session_id=item['session_id'], user_id=item['user_id'], level=item['level'], created_at=item['created_at'], expires_at=item['expires_at'], last_activity=item['last_activity'], permissions=item.get('permissions', ['*']) ) self.sessions[session.session_id] = session except Exception as e: print(f"加载会话失败: {e}") def _save_sessions(self): """保存会话""" data = [] for session in self.sessions.values(): data.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, "permissions": session.permissions }) try: with open(self.sessions_file, 'w', encoding='utf-8') as f: json.dump(data, f, ensure_ascii=False, indent=2) except Exception: pass def get_stats(self) -> dict: """获取认证统计""" active_sessions = sum( 1 for s in self.sessions.values() if time.time() < s.expires_at ) return { "active_sessions": active_sessions, "session_timeout": self.session_timeout } # === API认证装饰器 === def require_auth(required_level: str = "admin", permission: str = None): """ API认证装饰器 使用方式: @require_auth() def api_endpoint(self, handler): ... @require_auth(permission="sync:start") def api_sync_start(self, handler): ... """ def decorator(func): @wraps(func) def wrapper(self, handler, *args, **kwargs): # 检查是否需要认证 if required_level == "none": return func(self, handler, *args, **kwargs) config = getattr(handler, 'config', {}) auth_type = config.get('auth_type', 'none') # 如果auth_type为none,跳过认证 if auth_type == "none": handler.auth_result = { "authenticated": True, "level": "admin", "user_id": "anonymous", "permissions": ["*"] } return func(self, handler, *args, **kwargs) auth_manager = getattr(handler, 'auth_manager', None) if not auth_manager: handler.send_json_response({ "error": "认证系统未初始化", "code": "AUTH_NOT_INITIALIZED" }, 500) return auth_result = auth_manager.validate_request(handler, required_level) if not auth_result.get('authenticated'): handler.send_response(401) handler.send_header('WWW-Authenticate', 'Bearer realm="HYC API"') handler.send_header('Access-Control-Allow-Origin', '*') handler.send_json_response({ "error": "未认证或认证已过期", "code": "UNAUTHORIZED", "required_level": required_level, "auth_methods": [ "Authorization: Bearer ", "X-API-Key: ", f"Cookie: {auth_manager.cookie_name}=", "?key=" ] }) return if permission: if not auth_manager.check_permission(auth_result, permission): handler.send_json_response({ "error": "权限不足", "code": "FORBIDDEN", "required_permission": permission }, 403) return handler.auth_result = auth_result return func(self, handler, *args, **kwargs) return wrapper return decorator # === 需要认证的API端点定义 === # 规则格式: 'METHOD:/api/vN/path' 或 'METHOD:/api/vN/prefix/'(前缀规则,尾斜杠) # 匹配时路径统一归一化为无前导斜杠形式,与 v1/v2 传入的 api_action 对齐 ADMIN_API_ENDPOINTS = { # 同步管理 'POST:/api/v2/sync/': 'sync:manage', 'DELETE:/api/v2/sync/': 'sync:manage', # 缓存管理 'POST:/api/v2/cache/clean': 'cache:manage', 'POST:/api/v2/cache/prewarm/': 'cache:manage', 'DELETE:/api/v2/cache/prewarm/': 'cache:manage', # Webhook管理 'POST:/api/v2/webhooks': 'webhook:create', 'PUT:/api/v2/webhooks/': 'webhook:update', 'DELETE:/api/v2/webhooks/': 'webhook:delete', 'POST:/api/v2/webhooks/': 'webhook:trigger', # 服务器配置 'PUT:/api/v2/config': 'config:manage', 'POST:/api/v2/server/reload': 'server:reload', 'POST:/api/v2/server/': 'server:manage', # 文件管理(高危操作: 删除/重命名/元数据/版本) 'DELETE:/api/v1/file/': 'files:delete', 'DELETE:/api/v2/file/': 'files:delete', 'PUT:/api/v2/file/': 'files:update', 'POST:/api/v2/file/': 'files:update', # 用户管理 'POST:/api/v2/users': 'users:create', 'DELETE:/api/v2/users/': 'users:delete', 'PUT:/api/v2/users/': 'users:update', 'POST:/api/v2/user/password': 'users:update', # 镜像管理(写 settings.json,必须鉴权) 'POST:/api/v2/mirrors': 'mirrors:manage', 'PUT:/api/v2/mirrors/': 'mirrors:manage', 'DELETE:/api/v2/mirrors/': 'mirrors:manage', # 告警配置与API文档生成 'PUT:/api/v2/alerts': 'config:manage', 'POST:/api/v2/alerts': 'config:manage', 'POST:/api/v2/api-docs/generate': 'config:manage', # === API v1 文件操作认证 === 'PUT:/api/v1/mkdir': 'files:create', 'POST:/api/v1/upload': 'files:upload', 'POST:/api/v1/batch': 'files:batch', 'POST:/api/v1/archive': 'files:archive', } def _normalize_endpoint_patterns(): """把 ADMIN_API_ENDPOINTS 归一化为 (精确表, 前缀表) 规则格式: 'METHOD:/api/vN/path' - 无通配符 → 精确匹配 (method, path) - 尾斜杠或尾* → 前缀匹配, 匹配 norm_path.startswith(prefix + '/') """ exact = {} prefix = [] for pattern, permission in ADMIN_API_ENDPOINTS.items(): pat_method, pat_path = pattern.split(':', 1) pat_path = pat_path.lstrip('/') if '*' in pattern or pat_path.endswith('/'): if pat_path.endswith('*'): pat_path = pat_path.rstrip('*') pat_path = pat_path.rstrip('/') prefix.append((pat_method, pat_path, permission)) else: exact[(pat_method, pat_path)] = permission return exact, prefix _EXACT_ENDPOINTS, _PREFIX_ENDPOINTS = _normalize_endpoint_patterns() def check_endpoint_auth(method: str, path: str, auth_manager: APIAuthManager) -> dict: """检查端点是否需要认证 路径统一归一化为无前导斜杠形式(如 'api/v2/sync/sources'), 与 v1/v2 传入的 api_action 保持一致。 """ norm_path = path.lstrip('/') # 精确匹配 key = (method, norm_path) if key in _EXACT_ENDPOINTS: return { "required": True, "permission": _EXACT_ENDPOINTS[key] } # 前缀匹配 for pat_method, pat_path, permission in _PREFIX_ENDPOINTS: if method != pat_method and pat_method != '*': continue if norm_path.startswith(pat_path + '/'): return { "required": True, "permission": permission } return { "required": False, "permission": None }