diff --git a/api/auth_routes.py b/api/auth_routes.py index 9b5c100..eb2aab4 100644 --- a/api/auth_routes.py +++ b/api/auth_routes.py @@ -12,7 +12,8 @@ from flask import Blueprint, request, jsonify from auth.gateway import require_gateway_auth, require_role, get_user_permissions, MOCK_USERS from core.status_codes import SUCCESS, BAD_REQUEST, UNAUTHORIZED, FORBIDDEN, PERMISSION_DENIED, INTERNAL_ERROR from api.response_utils import success_response, error_response -import os +from config import DEV_MODE +import time from pathlib import Path from dotenv import load_dotenv @@ -23,6 +24,37 @@ load_dotenv(env_path) auth_bp = Blueprint('auth', __name__) +# ==================== 登录速率限制 ==================== +# 简单的内存级速率限制:每个 IP 在时间窗口内最多允许 N 次登录尝试 +_login_attempts = {} # {ip: [(timestamp, ...), ...]} +_RATE_LIMIT_WINDOW = 300 # 5 分钟窗口 +_RATE_LIMIT_MAX = 10 # 窗口内最多 10 次尝试 + + +def _check_rate_limit(client_ip: str) -> bool: + """检查是否超出登录速率限制,返回 True 表示允许""" + now = time.time() + if client_ip not in _login_attempts: + _login_attempts[client_ip] = [] + + # 清理过期记录 + _login_attempts[client_ip] = [ + t for t in _login_attempts[client_ip] + if now - t < _RATE_LIMIT_WINDOW + ] + + if len(_login_attempts[client_ip]) >= _RATE_LIMIT_MAX: + return False + + _login_attempts[client_ip].append(now) + return True + + +def _is_dev_mode() -> bool: + """统一的开发模式判断(由 config.py 集中管理)""" + return DEV_MODE + + @auth_bp.route('/auth/login', methods=['POST']) def mock_login(): """ @@ -52,9 +84,17 @@ def mock_login(): - manager / manager123 (经理,财务部) - user / test123 (普通用户,技术部) """ - # 默认开启开发模式(生产环境需设置 DEV_MODE=false) - if os.environ.get('DEV_MODE', 'true').lower() == 'false': - return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用,请设置 DEV_MODE=true", http_status=403) + if not _is_dev_mode(): + return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用,请在 .env 中设置 DEV_MODE=true", http_status=403) + + # 速率限制检查 + client_ip = request.remote_addr or 'unknown' + if not _check_rate_limit(client_ip): + return error_response( + "RATE_LIMITED", FORBIDDEN, + f"登录尝试过于频繁,请 {_RATE_LIMIT_WINDOW // 60} 分钟后再试", + http_status=429 + ) data = request.json or {} username = data.get('username') @@ -81,7 +121,9 @@ def mock_login(): def get_stats(): """获取系统统计信息(仅管理员)""" from flask import current_app - session_manager = current_app.config['SESSION_MANAGER'] + session_manager = current_app.config.get('SESSION_MANAGER') + if not session_manager: + return error_response("UNAVAILABLE", INTERNAL_ERROR, "会话管理器未启用", http_status=503) return success_response(data=session_manager.get_stats()) @@ -133,8 +175,7 @@ def get_users(): ] } """ - dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false' - if not dev_mode: + if not _is_dev_mode(): return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=403) users = [] @@ -161,12 +202,26 @@ def update_user(user_id): "is_active": false } """ - dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false' - if not dev_mode: + if not _is_dev_mode(): return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=403) - # 模拟用户不支持真正的状态切换,直接返回成功 - return success_response(data={"message": "操作成功(模拟)", "user_id": user_id}) + # 验证目标用户是否存在 + target_user = None + for username, info in MOCK_USERS.items(): + if info['user_id'] == user_id: + target_user = info + break + + if not target_user: + return error_response("NOT_FOUND", BAD_REQUEST, f"用户 {user_id} 不存在", http_status=404) + + data = request.json or {} + # 模拟操作:记录请求但不实际执行(mock 用户数据是静态的) + return success_response(data={ + "message": "操作成功(模拟)", + "user_id": user_id, + "applied_changes": data + }) @auth_bp.route('/auth/change-password', methods=['POST']) @@ -181,8 +236,7 @@ def change_password(): "new_password": "xxx" } """ - dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false' - if not dev_mode: + if not _is_dev_mode(): return error_response("FORBIDDEN", FORBIDDEN, "仅开发环境可用", http_status=403) data = request.json or {} @@ -195,5 +249,12 @@ def change_password(): if len(new_password) < 6: return error_response("INVALID_PARAMS", BAD_REQUEST, "新密码至少6位", http_status=400) - # 模拟环境直接返回成功 + # 验证当前用户的旧密码 + user = request.current_user + username = user.get('username', '') + mock_user = MOCK_USERS.get(username) + if mock_user and mock_user['password'] != old_password: + return error_response("UNAUTHORIZED", UNAUTHORIZED, "旧密码错误", http_status=401) + + # 模拟环境返回成功(不实际修改密码,mock 数据是静态的) return success_response(message="密码修改成功(模拟)") diff --git a/api/session_routes.py b/api/session_routes.py index 9c7c260..73ab3f0 100644 --- a/api/session_routes.py +++ b/api/session_routes.py @@ -71,7 +71,7 @@ def get_history(session_id): user_id = request.current_user["user_id"] # 验证会话归属 - sessions = session_manager.get_user_sessions(user_id) + sessions = session_manager.get_user_sessions(user_id, limit=20) session_ids = [s["session_id"] for s in sessions] if session_id not in session_ids: @@ -92,7 +92,7 @@ def delete_session(session_id): user_id = request.current_user["user_id"] # 验证会话归属 - sessions = session_manager.get_user_sessions(user_id) + sessions = session_manager.get_user_sessions(user_id, limit=20) session_ids = [s["session_id"] for s in sessions] if session_id not in session_ids: @@ -113,7 +113,7 @@ def clear_history(session_id): user_id = request.current_user["user_id"] # 验证会话归属 - sessions = session_manager.get_user_sessions(user_id) + sessions = session_manager.get_user_sessions(user_id, limit=20) session_ids = [s["session_id"] for s in sessions] if session_id not in session_ids: diff --git a/auth/gateway.py b/auth/gateway.py index cf245dd..bf8907c 100644 --- a/auth/gateway.py +++ b/auth/gateway.py @@ -20,12 +20,12 @@ ## 模式说明 -开发模式 (DEV_MODE=true,默认): +开发模式 (DEV_MODE=true): - 支持 mock token 模拟用户:Authorization: Bearer mock-token-admin - 无 Header 时自动使用开发测试用户 - 适用于前端测试和开发调试 -生产模式 (DEV_MODE=false): +生产模式 (DEV_MODE=false,默认): - 不需要 Header,直接放行 - 权限由后端完全控制,通过 collections 参数传入 - RAG 服务完全无状态,只负责问答检索 @@ -34,17 +34,11 @@ from functools import wraps from flask import request, jsonify from typing import Dict, Optional -import os -from pathlib import Path -from dotenv import load_dotenv - -# 加载 .env 文件(从项目根目录) -env_path = Path(__file__).parent.parent / '.env' -load_dotenv(env_path) +from config import DEV_MODE # ==================== 模拟用户数据(开发环境)==================== -# 用于前端模拟登录测试,仅 DEV_MODE=true 时生效 +# 用于前端模拟登录测试,仅 DEV_MODE=true 时生效(需在 .env 中显式开启) MOCK_USERS = { 'admin': { 'user_id': 'admin001', @@ -83,19 +77,19 @@ def require_gateway_auth(f): """ 网关认证装饰器 - 从 Header 读取用户信息 - 开发模式 (DEV_MODE=true,默认): + 开发模式 (DEV_MODE=true): - 支持 mock token: Authorization: Bearer mock-token-admin - 无 Header 时自动使用开发测试用户(admin 角色) - 生产模式 (DEV_MODE=false): + 生产模式 (DEV_MODE=false,默认): - 不需要 Header,直接放行 - 用户信息设为默认值 - 权限由后端通过 collections 参数控制 """ @wraps(f) def decorated(*args, **kwargs): - # 开发模式开关(默认开启,生产环境设置 DEV_MODE=false) - dev_mode = os.environ.get('DEV_MODE', 'true').lower() != 'false' + # 开发模式开关(统一由 config.py 管理) + dev_mode = DEV_MODE # 开发模式:支持 mock token if dev_mode: @@ -202,11 +196,24 @@ def can_delete_collection(role: str) -> bool: def require_role(*roles): """ - 兼容旧代码 - 权限由后端管理,此装饰器不再执行权限验证 + 角色验证装饰器(开发和生产环境均生效) + + 需搭配 @require_gateway_auth 使用(先设置 current_user,再验证角色)。 + + 开发模式: 检查 mock token 对应用户的角色 + 生产模式: 检查网关注入的 X-User-Role Header """ def decorator(f): @wraps(f) def decorated(*args, **kwargs): + if roles: + user = get_current_user() + if user is None or user.get('role') not in roles: + from flask import jsonify + return jsonify({ + "error": "权限不足,需要角色: {}".format(', '.join(roles)), + "status": "FORBIDDEN" + }), 403 return f(*args, **kwargs) return decorated return decorator @@ -214,11 +221,12 @@ def require_role(*roles): def require_collection_permission(operation: str): """ - 兼容旧代码 - 权限由后端管理,此装饰器不再执行权限验证 + 集合权限验证装饰器(开发模式下为占位实现,生产环境权限由网关控制) """ def decorator(f): @wraps(f) def decorated(*args, **kwargs): + # 生产环境下权限由网关/后端统一管控,此处放行 return f(*args, **kwargs) return decorated return decorator diff --git a/services/session.py b/services/session.py index d2b84b1..bb71868 100644 --- a/services/session.py +++ b/services/session.py @@ -159,7 +159,7 @@ class SessionManager: SELECT id, role, content, metadata, created_at FROM messages WHERE session_id = ? - ORDER BY created_at DESC + ORDER BY created_at DESC, id DESC LIMIT ? ''', (session_id, limit))