diff --git a/.env.example b/.env.example index aa93881..45fd0ca 100644 --- a/.env.example +++ b/.env.example @@ -79,7 +79,7 @@ BAODAN_WORKFLOW_API_KEY=app-your-workflow-api-key BAODAN_KB_API_KEY=dataset-your-kb-api-key # 访客模式 -GUEST_MODE=true +GUEST_MODE=false # 插件服务 PLUGIN_DAEMON_KEY=your_plugin_daemon_key diff --git a/api/insurance/app.py b/api/insurance/app.py index 93279e5..98135f4 100644 --- a/api/insurance/app.py +++ b/api/insurance/app.py @@ -124,7 +124,7 @@ def create_app() -> Flask: app.config['WECOM_WEBHOOK_URL'] = os.getenv('WECOM_WEBHOOK_URL', '') # 访客模式 - app.config['GUEST_MODE'] = os.getenv('GUEST_MODE', 'true').lower() == 'true' + app.config['GUEST_MODE'] = os.getenv('GUEST_MODE', 'false').lower() == 'true' # 域名(分享链接用) app.config['DOMAIN'] = os.getenv('DOMAIN', 'localhost') diff --git a/api/insurance/generation/routes.py b/api/insurance/generation/routes.py index 49cd2dc..dc41ac3 100644 --- a/api/insurance/generation/routes.py +++ b/api/insurance/generation/routes.py @@ -5,7 +5,7 @@ import json import logging from flask import Blueprint, request -from insurance.middleware.auth_middleware import jwt_required +from insurance.middleware.auth_middleware import account_required from insurance.utils.response import success, error, ErrorCode logger = logging.getLogger(__name__) @@ -16,7 +16,7 @@ workspace_bp = Blueprint("workspace", __name__) # ─── PPT 工作区 ──────────────────────────────────────────── @workspace_bp.route("/ppt/workspaces", methods=["GET"]) -@jwt_required +@account_required def list_ppt_workspaces(): """查询用户 PPT 工作区列表。""" from insurance.db.compat import db @@ -40,7 +40,7 @@ def list_ppt_workspaces(): @workspace_bp.route("/ppt/workspaces/", methods=["GET"]) -@jwt_required +@account_required def get_ppt_workspace(session_id): """获取 PPT 工作区详情。""" from insurance.models.ppt_session import PptSession @@ -53,7 +53,7 @@ def get_ppt_workspace(session_id): @workspace_bp.route("/ppt/workspaces//rename", methods=["PUT"]) -@jwt_required +@account_required def rename_ppt_workspace(session_id): """重命名 PPT 工作区。""" from insurance.db.compat import db @@ -75,7 +75,7 @@ def rename_ppt_workspace(session_id): @workspace_bp.route("/ppt/workspaces//draft", methods=["POST", "PATCH"]) -@jwt_required +@account_required def save_ppt_draft(session_id): """自动保存 PPT 草稿(乐观锁)。 @@ -116,7 +116,7 @@ def save_ppt_draft(session_id): @workspace_bp.route("/ppt/workspaces//archive", methods=["PUT"]) -@jwt_required +@account_required def archive_ppt_workspace(session_id): """归档 PPT 工作区。""" from datetime import datetime @@ -134,7 +134,7 @@ def archive_ppt_workspace(session_id): @workspace_bp.route("/ppt/workspaces//unarchive", methods=["PUT"]) -@jwt_required +@account_required def unarchive_ppt_workspace(session_id): """取消归档 PPT 工作区。""" from insurance.db.compat import db @@ -151,7 +151,7 @@ def unarchive_ppt_workspace(session_id): @workspace_bp.route("/ppt/workspaces//copy", methods=["POST"]) -@jwt_required +@account_required def copy_ppt_workspace(session_id): """复制 PPT 工作区为新任务。""" from insurance.db.compat import db @@ -183,7 +183,7 @@ def copy_ppt_workspace(session_id): # ─── 海报工作区 ──────────────────────────────────────────── @workspace_bp.route("/poster/workspaces", methods=["GET"]) -@jwt_required +@account_required def list_poster_workspaces(): """查询用户海报工作区列表。""" from insurance.models.poster_record import PosterRecord @@ -206,7 +206,7 @@ def list_poster_workspaces(): @workspace_bp.route("/poster/workspaces/", methods=["GET"]) -@jwt_required +@account_required def get_poster_workspace(record_id): """获取海报工作区详情。""" from insurance.models.poster_record import PosterRecord @@ -219,7 +219,7 @@ def get_poster_workspace(record_id): @workspace_bp.route("/poster/workspaces//rename", methods=["PUT"]) -@jwt_required +@account_required def rename_poster_workspace(record_id): """重命名海报工作区。""" from insurance.db.compat import db @@ -241,7 +241,7 @@ def rename_poster_workspace(record_id): @workspace_bp.route("/poster/workspaces//draft", methods=["POST", "PATCH"]) -@jwt_required +@account_required def save_poster_draft(record_id): """自动保存海报草稿(乐观锁)。""" from insurance.db.compat import db @@ -276,7 +276,7 @@ def save_poster_draft(record_id): @workspace_bp.route("/poster/workspaces//archive", methods=["PUT"]) -@jwt_required +@account_required def archive_poster_workspace(record_id): """归档海报工作区。""" from datetime import datetime @@ -295,7 +295,7 @@ def archive_poster_workspace(record_id): @workspace_bp.route("/poster/workspaces//copy", methods=["POST"]) -@jwt_required +@account_required def copy_poster_workspace(record_id): """复制海报工作区为新任务。""" from insurance.db.compat import db @@ -327,7 +327,7 @@ def copy_poster_workspace(record_id): # ─── 统一任务接口 ────────────────────────────────────────── @workspace_bp.route("/tasks", methods=["GET"]) -@jwt_required +@account_required def list_tasks(): """查询任务列表(任务中心)。""" from insurance.generation import task_service @@ -349,7 +349,7 @@ def list_tasks(): @workspace_bp.route("/tasks/active", methods=["GET"]) -@jwt_required +@account_required def list_active_tasks(): """查询活跃任务(任务坞)。""" from insurance.generation import task_service @@ -362,7 +362,7 @@ def list_active_tasks(): @workspace_bp.route("/tasks/", methods=["GET"]) -@jwt_required +@account_required def get_task(task_id): """获取任务详情。""" from insurance.generation import task_service @@ -375,7 +375,7 @@ def get_task(task_id): @workspace_bp.route("/tasks//cancel", methods=["PUT"]) -@jwt_required +@account_required def cancel_task(task_id): """取消排队中的任务。""" from insurance.generation import task_service @@ -388,7 +388,7 @@ def cancel_task(task_id): @workspace_bp.route("/tasks//hide", methods=["PUT"]) -@jwt_required +@account_required def hide_task(task_id): """从任务坞隐藏任务。""" from insurance.generation import task_service @@ -401,7 +401,7 @@ def hide_task(task_id): @workspace_bp.route("/tasks//viewed", methods=["PUT"]) -@jwt_required +@account_required def mark_task_viewed(task_id): """标记任务为已查看。""" from insurance.generation import task_service @@ -414,7 +414,7 @@ def mark_task_viewed(task_id): @workspace_bp.route("/tasks//download", methods=["GET"]) -@jwt_required +@account_required def download_task_output(task_id): """按任务下载成品(确保下载的是该任务的版本,而非工作区最新版本)。""" import json as _json diff --git a/api/insurance/generation/task_service.py b/api/insurance/generation/task_service.py index e1b960f..35d6909 100644 --- a/api/insurance/generation/task_service.py +++ b/api/insurance/generation/task_service.py @@ -21,6 +21,7 @@ def create_task(user_id: str, artifact_type: str, operation: str, # 幂等检查 if idempotency_key: existing = GenerationTask.query.filter_by( + user_id=user_id, idempotency_key=idempotency_key, ).filter(GenerationTask.status.in_(["queued", "running", "done"])).first() if existing: @@ -28,6 +29,9 @@ def create_task(user_id: str, artifact_type: str, operation: str, # 检查同一工作区是否有运行中的任务 active = GenerationTask.query.filter_by( + user_id=user_id, + artifact_type=artifact_type, + operation=operation, workspace_id=workspace_id, ).filter(GenerationTask.status.in_(["queued", "running"])).first() if active: @@ -104,7 +108,7 @@ def _sync_failed_ppt_session(task, message: str): from insurance.models.ppt_session import PptSession session = PptSession.query.get(task.workspace_id) - if not session or session.status != "parsing": + if not session or session.user_id != task.user_id or session.status != "parsing": return if session.latest_task_id and session.latest_task_id != task.id: return @@ -138,7 +142,7 @@ def _sync_ppt_workspace(task, db): from insurance.models.ppt_session import PptSession session = PptSession.query.get(task.workspace_id) - if not session: + if not session or session.user_id != task.user_id: return # 安全检查:只同步最新任务 @@ -180,7 +184,7 @@ def _sync_poster_workspace(task, db): from insurance.models.poster_record import PosterRecord record = PosterRecord.query.get(task.workspace_id) - if not record: + if not record or record.user_id != task.user_id: return # 安全检查 diff --git a/api/insurance/middleware/auth_middleware.py b/api/insurance/middleware/auth_middleware.py index c2d7d89..d18ae83 100644 --- a/api/insurance/middleware/auth_middleware.py +++ b/api/insurance/middleware/auth_middleware.py @@ -93,6 +93,18 @@ def jwt_required(f): return decorated +def account_required(f): + """要求已登录账号,不允许访客身份访问私有业务数据。""" + @wraps(f) + @jwt_required + def decorated(*args, **kwargs): + user_id = str(getattr(request, "user_id", "")) + if not user_id or user_id.startswith("guest_"): + return jsonify({"code": 1002, "message": "请先登录账号", "data": None}), 401 + return f(*args, **kwargs) + return decorated + + def admin_required(f): """装饰器:要求管理员或超级管理员权限。""" @wraps(f) diff --git a/api/insurance/models/generation_task.py b/api/insurance/models/generation_task.py index 2e8c2a4..3832c5a 100644 --- a/api/insurance/models/generation_task.py +++ b/api/insurance/models/generation_task.py @@ -79,6 +79,7 @@ class GenerationTask(db.Model): import json if idempotency_key: existing = GenerationTask.query.filter_by( + user_id=user_id, idempotency_key=idempotency_key, ).filter(GenerationTask.status.in_(["queued", "running", "done"])).first() if existing: diff --git a/api/insurance/poster/routes.py b/api/insurance/poster/routes.py index 2bb2924..df19117 100644 --- a/api/insurance/poster/routes.py +++ b/api/insurance/poster/routes.py @@ -1,7 +1,7 @@ """海报功能路由。""" import os from flask import Blueprint, request, jsonify, send_file -from insurance.middleware.auth_middleware import jwt_required +from insurance.middleware.auth_middleware import account_required as jwt_required from insurance.utils.response import success, error, ErrorCode poster_bp = Blueprint("poster", __name__) diff --git a/api/insurance/ppt/routes.py b/api/insurance/ppt/routes.py index 323220c..4f3a951 100644 --- a/api/insurance/ppt/routes.py +++ b/api/insurance/ppt/routes.py @@ -5,7 +5,7 @@ import json import logging from datetime import datetime from flask import Blueprint, request, jsonify, send_file -from insurance.middleware.auth_middleware import jwt_required +from insurance.middleware.auth_middleware import account_required as jwt_required from insurance.utils.response import success, error, ErrorCode logger = logging.getLogger(__name__) diff --git a/frontend/src/utils/api.ts b/frontend/src/utils/api.ts index 6da1cc1..6fb12c3 100644 --- a/frontend/src/utils/api.ts +++ b/frontend/src/utils/api.ts @@ -15,13 +15,7 @@ export function getAuthHeaders(): Record { if (token) { return { Authorization: `Bearer ${token}` } } - // 访客模式:使用固定的 guest ID - let guestId = localStorage.getItem('guest_id') - if (!guestId) { - guestId = 'guest_' + Math.random().toString(36).substring(2, 10) - localStorage.setItem('guest_id', guestId) - } - return { Authorization: `Bearer guest_${guestId}` } + return {} } // 请求拦截器:自动添加 JWT @@ -31,13 +25,7 @@ api.interceptors.request.use((config) => { if (token) { config.headers.Authorization = `Bearer ${token}` } else { - // 访客模式:使用固定的 guest ID - let guestId = localStorage.getItem('guest_id') - if (!guestId) { - guestId = 'guest_' + Math.random().toString(36).substring(2, 10) - localStorage.setItem('guest_id', guestId) - } - config.headers.Authorization = `Bearer guest_${guestId}` + delete config.headers.Authorization } return config }) diff --git a/tests/generation_routes_test.py b/tests/generation_routes_test.py index 1a9805e..bc8d9f6 100644 --- a/tests/generation_routes_test.py +++ b/tests/generation_routes_test.py @@ -3,6 +3,7 @@ from pathlib import Path import sys from flask import Flask +import jwt sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "api")) @@ -21,6 +22,34 @@ def test_task_detail_returns_task_as_direct_data(monkeypatch): }, ) + app = Flask(__name__) + app.config["JWT_SECRET"] = "test-secret" + app.config["GUEST_MODE"] = True + app.register_blueprint(workspace_bp, url_prefix="/insurance/workspace") + token = jwt.encode({"user_id": "user-1"}, "test-secret", algorithm="HS256") + + response = app.test_client().get( + "/insurance/workspace/tasks/task-1", + headers={"Authorization": f"Bearer {token}"}, + ) + + assert response.status_code == 200 + body = response.get_json() + assert body["data"]["status"] == "running" + assert "data" not in body["data"] + + +def test_generation_tasks_reject_guest_identity(monkeypatch): + """生成物是账号私有数据,访客令牌不能读取任务。""" + called = False + + def fake_get_task(_task_id, _user_id): + nonlocal called + called = True + return {"code": 0, "data": {}} + + monkeypatch.setattr(task_service, "get_task", fake_get_task) + app = Flask(__name__) app.config["GUEST_MODE"] = True app.register_blueprint(workspace_bp, url_prefix="/insurance/workspace") @@ -30,7 +59,6 @@ def test_task_detail_returns_task_as_direct_data(monkeypatch): headers={"Authorization": "Bearer guest_test"}, ) - assert response.status_code == 200 - body = response.get_json() - assert body["data"]["status"] == "running" - assert "data" not in body["data"] + assert response.status_code == 401 + assert response.get_json()["message"] == "请先登录账号" + assert called is False diff --git a/tests/ppt_task_lifecycle_test.py b/tests/ppt_task_lifecycle_test.py index 2c0f536..4001165 100644 --- a/tests/ppt_task_lifecycle_test.py +++ b/tests/ppt_task_lifecycle_test.py @@ -105,6 +105,50 @@ def test_dispatch_failure_finishes_task(monkeypatch): assert synced and synced[0][0] == result["data"]["id"] +def test_create_task_idempotency_is_scoped_to_user(monkeypatch): + """相同幂等键不能让一个用户复用另一个用户的任务。""" + filter_by_calls = [] + created_tasks = [] + other_user_task = _GenerationTask(user_id="user-1") + other_user_task.status = "done" + + class OwnerAwareQuery(_Query): + def filter_by(self, **kwargs): + filter_by_calls.append(kwargs) + return _Query() if kwargs.get("user_id") == "user-2" else self + + def first(self): + return other_user_task + + class OwnerScopedGenerationTask(_GenerationTask): + query = OwnerAwareQuery() + + def __init__(self, **kwargs): + super().__init__(**kwargs) + created_tasks.append(kwargs) + + fake_db = types.SimpleNamespace(session=_Session()) + compat_module = types.ModuleType("insurance.db.compat") + compat_module.db = fake_db + model_module = types.ModuleType("insurance.models.generation_task") + model_module.GenerationTask = OwnerScopedGenerationTask + monkeypatch.setitem(sys.modules, "insurance.db.compat", compat_module) + monkeypatch.setitem(sys.modules, "insurance.models.generation_task", model_module) + monkeypatch.setattr(task_service, "_dispatch_to_celery", lambda _task: None) + + result = task_service.create_task( + user_id="user-2", + artifact_type="ppt", + operation="generate", + workspace_id="session-2", + idempotency_key="shared-key", + ) + + assert result["code"] == 0 + assert created_tasks[0]["user_id"] == "user-2" + assert {"user_id": "user-2", "idempotency_key": "shared-key"} in filter_by_calls + + def test_cached_extraction_reports_completion(tmp_path, monkeypatch): """重复解析应命中缓存,并向调用方报告完成阶段。""" pdf_path = tmp_path / "plan.pdf"