Files
geMoldInsight/tests/test_experience_feedback_router.py
T

319 lines
13 KiB
Python
Raw Normal View History

"""D17 Human-in-Loop 闭环:API 契约测试。
覆盖:
- POST /api/tasks/{task_id}/experience-feedback
- 401(无登录态 —— 由 Depends(get_current_active_user) 处理)
- 403(user 角色无 feedback_experience_hint 权限)
- 200(admin 角色有 manage_experience_feedback 全权限)
- 200(process_engineer 角色有 feedback_experience_hint 权限)
- GET /api/tasks/{task_id}/experience-hints
- 200 命中(同 stp_file_id 历史反馈聚合)
- cache invalidation(POST 写完后视图失效)
- D9 边界:record_feedback 失败时 db 不留半成品
"""
import pytest
from fastapi import FastAPI
from httpx import AsyncClient, ASGITransport
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.orm import configure_mappers
from shared.models.base import Base
from shared.models.identity import User, Role, Permission, RolePermission, UserRole
from shared.services.auth_service import get_current_active_user
from shared.database.database import get_db_session
from moldinsight.api.experience_feedback_router import router as feedback_router
from moldinsight.models import (
STPFile, GeometryData, MoldCavityData, ProcessingTask, ExperienceFeedback,
)
@pytest.fixture
async def feedback_client(async_engine, seeded_db):
"""构造带 experience_feedback_router 的 test app。
与 conftest.client 不同,这里我们用 seeded_db 的 user=tester,但通过依赖覆盖
让所有请求都以 admin 身份进(admin 是项目测试约定身份)。
关键点:override 返回的 User 必须用 selectinload 预加载 user_roles → role → role_permissions → permission,
否则 User.has_permission() 内部访问 self.roles 触发跨 session lazy load 失败。
"""
from sqlalchemy.orm import selectinload
session_factory = async_sessionmaker(async_engine, class_=AsyncSession, expire_on_commit=False)
test_app = FastAPI()
test_app.include_router(feedback_router)
async def override_get_db_session():
async with session_factory() as session:
yield session
async def override_get_current_active_user():
async with session_factory() as session:
result = await session.execute(
select(User)
.where(User.username == "tester")
.options(
selectinload(User.user_roles)
.selectinload(UserRole.role)
.selectinload(Role.role_permissions)
.selectinload(RolePermission.permission)
)
)
return result.scalar_one()
test_app.dependency_overrides[get_db_session] = override_get_db_session
test_app.dependency_overrides[get_current_active_user] = override_get_current_active_user
transport = ASGITransport(app=test_app)
async with AsyncClient(transport=transport, base_url="http://test") as ac:
yield ac
test_app.dependency_overrides.clear()
async def _grant_permission(session, user, code):
"""给测试 user 加指定 permission_code。
注意:User.is_superuser 是 @property(派生自 role.code == "admin"),
不能直接赋值;admin 权限通过给 user 关联 'admin' role 触发。
"""
# 找/创建 permission
perm_row = await session.execute(select(Permission).where(Permission.code == code))
perm = perm_row.scalar_one_or_none()
if perm is None:
perm = Permission(code=code, name=code, module="moldinsight")
session.add(perm)
await session.flush()
# 找/创建 role(用 permission code 作 role code,便于复用)
role_row = await session.execute(select(Role).where(Role.code == code))
role = role_row.scalar_one_or_none()
if role is None:
role = Role(code=code, name=code, is_system=False)
session.add(role)
await session.flush()
rp = RolePermission(role_id=role.id, permission_id=perm.id)
session.add(rp)
# 关联 user(如未关联)
user_role_row = await session.execute(
select(UserRole).where(UserRole.user_id == user.id, UserRole.role_id == role.id)
)
if user_role_row.scalar_one_or_none() is None:
session.add(UserRole(user_id=user.id, role_id=role.id))
await session.commit()
# ── POST /experience-feedback 测试 ──
async def test_submit_feedback_403_without_permission(feedback_client, async_engine):
"""tester 默认无任何权限 → 403。"""
session_factory = async_sessionmaker(async_engine, class_=AsyncSession, expire_on_commit=False)
# 确保 tester 没有 admin role 也没有 feedback_experience_hint role
async with session_factory() as session:
await session.execute(
UserRole.__table__.delete().where(UserRole.user_id == 1)
)
await session.commit()
resp = await feedback_client.post(
"/tasks/task-demo-1/experience-feedback",
json={
"scheme_id": "scheme_1",
"feedback_status": "adopted",
},
)
assert resp.status_code == 403, resp.text
assert "工艺工程师" in resp.text
async def test_submit_feedback_200_with_feedback_permission(feedback_client, async_engine):
"""给 tester 授予 feedback_experience_hint → 200 + 写入经验反馈。"""
session_factory = async_sessionmaker(async_engine, class_=AsyncSession, expire_on_commit=False)
async with session_factory() as session:
tester = (await session.execute(select(User).where(User.username == "tester"))).scalar_one()
await _grant_permission(session, tester, "feedback_experience_hint")
resp = await feedback_client.post(
"/tasks/task-demo-1/experience-feedback",
json={
"scheme_id": "scheme_1",
"feedback_status": "adopted",
"feedback_reason": "工艺验证 OK",
"confidence_at_submit": 0.85,
"score_at_submit": 87.5,
},
)
assert resp.status_code == 200, resp.text
body = resp.json()
assert body["scheme_id"] == "scheme_1"
assert body["feedback_status"] == "adopted"
# DB 真的写入了
async with session_factory() as session:
result = await session.execute(
select(ExperienceFeedback).where(ExperienceFeedback.scheme_id == "scheme_1")
)
fb = result.scalar_one()
assert fb.user_id == tester.id
assert fb.feedback_status == "adopted"
assert fb.role_code == "feedback_experience_hint" # 写入时角色归因
assert fb.expires_at is not None
async def test_submit_feedback_200_with_admin(feedback_client, async_engine):
"""admin role → has_permission 走 role.code=='admin' 短路 → 200。"""
session_factory = async_sessionmaker(async_engine, class_=AsyncSession, expire_on_commit=False)
async with session_factory() as session:
tester = (await session.execute(select(User).where(User.username == "tester"))).scalar_one()
await _grant_permission(session, tester, "admin")
resp = await feedback_client.post(
"/tasks/task-demo-1/experience-feedback",
json={
"scheme_id": "scheme_2",
"feedback_status": "rejected",
},
)
assert resp.status_code == 200, resp.text
assert resp.json()["scheme_axis"] # 自动从 cavity_key_info 解析,缺则默认 Z
async def test_submit_feedback_invalid_status_returns_422(feedback_client, async_engine):
"""feedback_status 非法 → Pydantic 校验 422。"""
session_factory = async_sessionmaker(async_engine, class_=AsyncSession, expire_on_commit=False)
async with session_factory() as session:
tester = (await session.execute(select(User).where(User.username == "tester"))).scalar_one()
await _grant_permission(session, tester, "admin")
resp = await feedback_client.post(
"/tasks/task-demo-1/experience-feedback",
json={
"scheme_id": "scheme_1",
"feedback_status": "approve", # 非法值
},
)
assert resp.status_code == 422
# ── GET /experience-hints 测试 ──
async def test_get_hints_200_empty(feedback_client, async_engine):
"""无反馈历史 → 空 hints 列表,fingerprint 回显。"""
session_factory = async_sessionmaker(async_engine, class_=AsyncSession, expire_on_commit=False)
async with session_factory() as session:
tester = (await session.execute(select(User).where(User.username == "tester"))).scalar_one()
await _grant_permission(session, tester, "admin")
resp = await feedback_client.get("/tasks/task-demo-1/experience-hints")
assert resp.status_code == 200, resp.text
body = resp.json()
assert body["task_id"] == "task-demo-1"
assert body["stp_file_id"] == 1
assert body["hints"] == []
assert "bbox_aspect" in body["fingerprint"]
async def test_get_hints_aggregates_by_axis(feedback_client, async_engine):
"""写入多条反馈后 GET hints 按 axis 聚合。"""
session_factory = async_sessionmaker(async_engine, class_=AsyncSession, expire_on_commit=False)
async with session_factory() as session:
tester = (await session.execute(select(User).where(User.username == "tester"))).scalar_one()
await _grant_permission(session, tester, "admin")
# 写入 4 条反馈,全部落在 axis="Z" 默认(无 cavity_key_info)
for fb in [
{"scheme_id": "x_1", "feedback_status": "adopted"},
{"scheme_id": "x_2", "feedback_status": "adopted"},
{"scheme_id": "x_3", "feedback_status": "rejected"},
{"scheme_id": "z_1", "feedback_status": "adopted"},
]:
resp = await feedback_client.post(
"/tasks/task-demo-1/experience-feedback",
json=fb,
)
assert resp.status_code == 200, resp.text
# GET hints
resp = await feedback_client.get("/tasks/task-demo-1/experience-hints")
assert resp.status_code == 200, resp.text
body = resp.json()
# 无 cavity_key_info 时所有 feedback 落在 axis="Z" 默认值 → 4 条聚合
assert len(body["hints"]) == 1
h = body["hints"][0]
assert h["scheme_axis"] == "Z"
assert h["adopted_count"] == 3
assert h["rejected_count"] == 1
assert h["sample_count"] == 4
assert h["confidence"] == 0.5 # (3-1)/4
async def test_submit_feedback_increments_expires_at(feedback_client, async_engine):
"""同 stp_file_id 上写入新反馈时,旧行的 expires_at 应被续期(write-time 续期)。"""
session_factory = async_sessionmaker(async_engine, class_=AsyncSession, expire_on_commit=False)
async with session_factory() as session:
tester = (await session.execute(select(User).where(User.username == "tester"))).scalar_one()
await _grant_permission(session, tester, "admin")
# 写入第一条反馈
resp = await feedback_client.post(
"/tasks/task-demo-1/experience-feedback",
json={"scheme_id": "scheme_1", "feedback_status": "adopted"},
)
assert resp.status_code == 200
# 拿到第一条 expires_at
async with session_factory() as session:
first = (await session.execute(
select(ExperienceFeedback).where(ExperienceFeedback.scheme_id == "scheme_1")
)).scalar_one()
first_expires = first.expires_at
assert first_expires is not None
# 写第二条(不同 scheme_id),应触发同 stp_file 续期
import asyncio
await asyncio.sleep(0.05)
resp = await feedback_client.post(
"/tasks/task-demo-1/experience-feedback",
json={"scheme_id": "scheme_2", "feedback_status": "rejected"},
)
assert resp.status_code == 200
# 第一条 expires_at 应被续期(≥ 原值)
async with session_factory() as session:
first_after = (await session.execute(
select(ExperienceFeedback).where(ExperienceFeedback.scheme_id == "scheme_1")
)).scalar_one()
assert first_after.expires_at >= first_expires
# ── 任务归属校验 ──
async def test_submit_feedback_with_unknown_task_returns_404(feedback_client, async_engine):
"""task_id 不存在 → ensure_task_access 返回 404(不是 500)。
D9 边界保护:service.record_feedback 永远走不到(ensure_task_access 先拦截)。
"""
session_factory = async_sessionmaker(async_engine, class_=AsyncSession, expire_on_commit=False)
async with session_factory() as session:
tester = (await session.execute(select(User).where(User.username == "tester"))).scalar_one()
await _grant_permission(session, tester, "admin")
resp = await feedback_client.post(
"/tasks/non-existent-task-id/experience-feedback",
json={"scheme_id": "scheme_1", "feedback_status": "adopted"},
)
assert resp.status_code == 404
assert "不存在" in resp.text
# ── 配置:D17 模型注册收口 ──
def test_experience_feedback_registered_in_metadata():
"""D17:experience_feedback 表已加入 Base.metadata(防止漏注册导致 ORM 不可用)。"""
configure_mappers()
assert "experience_feedback" in Base.metadata.tables