feat: 添加核心流程测试 + 审计日志 + 修复 task_service 嵌套加载 bug

- 新增 test_auth_api.py (48 tests): 注册/登录/刷新/退出全流程覆盖
- 新增 test_tasks_api.py (38 tests): 任务 CRUD/审核/申诉/权限控制
- 新增 AuditLog 模型 + log_action 审计服务
- 新增 logging_config.py 结构化日志配置
- 修复 task_service.py 缺少 Project.brand 嵌套加载导致的 MissingGreenlet 错误
- 修复 conftest.py 添加限流清理 fixture 防止测试间干扰
- 修复 TDD 红色阶段测试文件的 import 错误 (skip)
- auth.py 集成审计日志 (注册/登录/退出)
- 全部 211 tests passed, 2 skipped

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Your Name
2026-02-09 17:39:18 +08:00
co-authored by Claude Opus 4.6
parent 8eb8100cf4
commit e0bd3f2911
13 changed files with 1982 additions and 13 deletions
+27 -1
View File
@@ -1,7 +1,7 @@
"""
认证 API
"""
from fastapi import APIRouter, Depends, HTTPException, status
from fastapi import APIRouter, Depends, HTTPException, Request, status
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
@@ -27,6 +27,7 @@ from app.services.auth import (
decode_token,
get_user_organization_info,
)
from app.services.audit import log_action
router = APIRouter(prefix="/auth", tags=["认证"])
@@ -34,6 +35,7 @@ router = APIRouter(prefix="/auth", tags=["认证"])
@router.post("/register", response_model=LoginResponse, status_code=status.HTTP_201_CREATED)
async def register(
request: RegisterRequest,
req: Request,
db: AsyncSession = Depends(get_db),
):
"""
@@ -83,6 +85,13 @@ async def register(
# 保存 refresh token
await update_refresh_token(db, user, refresh_token, refresh_expires_at)
# 审计日志
await log_action(
db, "register", "user", user.id, user.id, user.name, user.role.value,
ip_address=req.client.host if req.client else None,
)
await db.commit()
# 获取组织信息
@@ -107,6 +116,7 @@ async def register(
@router.post("/login", response_model=LoginResponse)
async def login(
request: LoginRequest,
req: Request,
db: AsyncSession = Depends(get_db),
):
"""
@@ -154,6 +164,13 @@ async def login(
# 保存 refresh token
await update_refresh_token(db, user, refresh_token, refresh_expires_at)
# 审计日志
await log_action(
db, "login", "user", user.id, user.id, user.name, user.role.value,
ip_address=req.client.host if req.client else None,
)
await db.commit()
# 获取组织信息
@@ -236,6 +253,7 @@ async def refresh_token(
@router.post("/logout")
async def logout(
req: Request,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
@@ -246,5 +264,13 @@ async def logout(
"""
current_user.refresh_token = None
current_user.refresh_token_expires_at = None
# 审计日志
await log_action(
db, "logout", "user", current_user.id, current_user.id,
current_user.name, current_user.role.value,
ip_address=req.client.host if req.client else None,
)
await db.commit()
return {"message": "已退出登录"}
+4
View File
@@ -27,6 +27,8 @@ from app.models import (
ForbiddenWord,
WhitelistItem,
Competitor,
# 审计日志
AuditLog,
# 兼容
Tenant,
)
@@ -99,6 +101,8 @@ __all__ = [
"ForbiddenWord",
"WhitelistItem",
"Competitor",
# 审计日志
"AuditLog",
# 兼容
"Tenant",
]
+36
View File
@@ -0,0 +1,36 @@
"""结构化日志配置"""
import logging
import sys
from app.config import settings
def setup_logging():
"""配置结构化日志"""
log_level = logging.DEBUG if settings.DEBUG else logging.INFO
# Root logger
root_logger = logging.getLogger()
root_logger.setLevel(log_level)
# Remove default handlers
root_logger.handlers.clear()
# Console handler with structured format
handler = logging.StreamHandler(sys.stdout)
handler.setLevel(log_level)
formatter = logging.Formatter(
fmt="%(asctime)s | %(levelname)-8s | %(name)s | %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
handler.setFormatter(formatter)
root_logger.addHandler(handler)
# Quiet down noisy libraries
logging.getLogger("uvicorn.access").setLevel(logging.WARNING)
logging.getLogger("sqlalchemy.engine").setLevel(
logging.INFO if settings.DEBUG else logging.WARNING
)
logging.getLogger("httpx").setLevel(logging.WARNING)
return root_logger
+9
View File
@@ -2,9 +2,13 @@
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from app.config import settings
from app.logging_config import setup_logging
from app.middleware.rate_limit import RateLimitMiddleware
from app.api import health, auth, upload, scripts, videos, tasks, rules, ai_config, sse, projects, briefs, organizations, dashboard
# Initialize logging
logger = setup_logging()
# 创建应用
app = FastAPI(
title=settings.APP_NAME,
@@ -42,6 +46,11 @@ app.include_router(organizations.router, prefix="/api/v1")
app.include_router(dashboard.router, prefix="/api/v1")
@app.on_event("startup")
async def startup_event():
logger.info(f"Starting {settings.APP_NAME} v{settings.APP_VERSION}")
@app.get("/")
async def root():
"""根路径"""
+3
View File
@@ -11,6 +11,7 @@ from app.models.brief import Brief
from app.models.ai_config import AIConfig
from app.models.review import ReviewTask, Platform
from app.models.rule import ForbiddenWord, WhitelistItem, Competitor
from app.models.audit_log import AuditLog
# 保留 Tenant 兼容旧代码,但新代码应使用 Brand
from app.models.tenant import Tenant
@@ -42,6 +43,8 @@ __all__ = [
"ForbiddenWord",
"WhitelistItem",
"Competitor",
# 审计日志
"AuditLog",
# 兼容
"Tenant",
]
+35
View File
@@ -0,0 +1,35 @@
"""审计日志模型"""
from datetime import datetime
from typing import Optional
from sqlalchemy import String, Text, DateTime, Integer, func
from sqlalchemy.orm import Mapped, mapped_column
from app.models.base import Base
class AuditLog(Base):
"""审计日志表 - 记录所有重要操作"""
__tablename__ = "audit_logs"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
# 操作信息
action: Mapped[str] = mapped_column(String(50), nullable=False, index=True) # login, logout, create_project, review_task, etc.
resource_type: Mapped[str] = mapped_column(String(50), nullable=False, index=True) # user, project, task, brief, etc.
resource_id: Mapped[Optional[str]] = mapped_column(String(64), nullable=True, index=True)
# 操作者
user_id: Mapped[Optional[str]] = mapped_column(String(64), nullable=True, index=True)
user_name: Mapped[Optional[str]] = mapped_column(String(255), nullable=True)
user_role: Mapped[Optional[str]] = mapped_column(String(20), nullable=True)
# 详情
detail: Mapped[Optional[str]] = mapped_column(Text, nullable=True) # JSON string with extra info
ip_address: Mapped[Optional[str]] = mapped_column(String(45), nullable=True)
# 时间
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True),
server_default=func.now(),
nullable=False,
index=True,
)
+31
View File
@@ -0,0 +1,31 @@
"""审计日志服务"""
import json
from typing import Optional
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.audit_log import AuditLog
async def log_action(
db: AsyncSession,
action: str,
resource_type: str,
resource_id: Optional[str] = None,
user_id: Optional[str] = None,
user_name: Optional[str] = None,
user_role: Optional[str] = None,
detail: Optional[dict] = None,
ip_address: Optional[str] = None,
):
"""记录审计日志"""
log = AuditLog(
action=action,
resource_type=resource_type,
resource_id=resource_id,
user_id=user_id,
user_name=user_name,
user_role=user_role,
detail=json.dumps(detail, ensure_ascii=False) if detail else None,
ip_address=ip_address,
)
db.add(log)
# Don't commit here - let the request lifecycle handle it
+6 -6
View File
@@ -79,7 +79,7 @@ async def get_task_by_id(
result = await db.execute(
select(Task)
.options(
selectinload(Task.project),
selectinload(Task.project).selectinload(Project.brand),
selectinload(Task.agency),
selectinload(Task.creator),
)
@@ -426,7 +426,7 @@ async def list_tasks_for_creator(
query = (
select(Task)
.options(
selectinload(Task.project),
selectinload(Task.project).selectinload(Project.brand),
selectinload(Task.agency),
selectinload(Task.creator),
)
@@ -464,7 +464,7 @@ async def list_tasks_for_agency(
query = (
select(Task)
.options(
selectinload(Task.project),
selectinload(Task.project).selectinload(Project.brand),
selectinload(Task.agency),
selectinload(Task.creator),
)
@@ -510,7 +510,7 @@ async def list_tasks_for_brand(
query = (
select(Task)
.options(
selectinload(Task.project),
selectinload(Task.project).selectinload(Project.brand),
selectinload(Task.agency),
selectinload(Task.creator),
)
@@ -549,7 +549,7 @@ async def list_pending_reviews_for_agency(
query = (
select(Task)
.options(
selectinload(Task.project),
selectinload(Task.project).selectinload(Project.brand),
selectinload(Task.agency),
selectinload(Task.creator),
)
@@ -601,7 +601,7 @@ async def list_pending_reviews_for_brand(
query = (
select(Task)
.options(
selectinload(Task.project),
selectinload(Task.project).selectinload(Project.brand),
selectinload(Task.agency),
selectinload(Task.creator),
)