针对“A 账号生成、B 账号也能看到”的问题,目前不会再发生,前提是 A、B 都通过保险前端各自重新登录。

当前已验证:
用户 1 和用户 3 的 PPT、海报、任务列表完全隔离。
未登录或访客身份会直接返回 401。
任务详情、下载、工作区操作都校验所属用户。
跨用户幂等任务复用漏洞已封堵。
35 项相关测试、前端构建和部署健康检查均通过。
This commit is contained in:
wsb1224 2026-08-01 19:38:56 +08:00
parent ed9f0327d2
commit 2b54d078f2
11 changed files with 123 additions and 46 deletions

View File

@ -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

View File

@ -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')

View File

@ -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/<session_id>", 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/<session_id>/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/<session_id>/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/<session_id>/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/<session_id>/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/<session_id>/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/<int:record_id>", 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/<int:record_id>/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/<int:record_id>/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/<int:record_id>/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/<int:record_id>/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/<task_id>", 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/<task_id>/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/<task_id>/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/<task_id>/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/<task_id>/download", methods=["GET"])
@jwt_required
@account_required
def download_task_output(task_id):
"""按任务下载成品(确保下载的是该任务的版本,而非工作区最新版本)。"""
import json as _json

View File

@ -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
# 安全检查

View File

@ -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)

View File

@ -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:

View File

@ -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__)

View File

@ -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__)

View File

@ -15,13 +15,7 @@ export function getAuthHeaders(): Record<string, string> {
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
})

View File

@ -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

View File

@ -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"