style: normalize python formatting gates

This commit is contained in:
Xiaoxia AI
2026-06-21 06:52:19 +08:00
parent 0809a079c5
commit bfbaddbd9a
129 changed files with 3024 additions and 2485 deletions
+1 -2
View File
@@ -1,5 +1,3 @@
from fastapi import APIRouter
from app.api.routes.asset_libraries import router as asset_libraries_router
from app.api.routes.assets import router as assets_router
from app.api.routes.auth_simple import router as auth_router
@@ -11,6 +9,7 @@ from app.api.routes.ingest_jobs import router as ingest_jobs_router
from app.api.routes.project_management import router as project_management_router
from app.api.routes.projects import router as projects_router
from app.api.routes.upload import router as upload_router
from fastapi import APIRouter
api_router = APIRouter(prefix="/api/v1")
health_router = APIRouter()
+13 -4
View File
@@ -1,9 +1,18 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.dependencies import get_asset_library_repository
from app.schemas.asset_library import AssetLibraryResponse, CreateAssetLibraryRequest, ListAssetLibrariesResponse
from typing import Any
from packages.application import CreateAssetLibraryCommand, CreateAssetLibraryUseCase, ListAssetLibrariesUseCase
from app.schemas.asset_library import (
AssetLibraryResponse,
CreateAssetLibraryRequest,
ListAssetLibrariesResponse,
)
from fastapi import APIRouter, Depends
from packages.application import (
CreateAssetLibraryCommand,
CreateAssetLibraryUseCase,
ListAssetLibrariesUseCase,
)
from packages.domain import AssetLibraryKind
router = APIRouter()
+8 -3
View File
@@ -1,9 +1,14 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.dependencies import get_asset_repository
from app.schemas.asset import AssetResponse, CreateAssetRequest, ListAssetsResponse
from typing import Any
from packages.application import CreateAssetCommand, CreateAssetUseCase, ListAssetsUseCase
from fastapi import APIRouter, Depends
from packages.application import (
CreateAssetCommand,
CreateAssetUseCase,
ListAssetsUseCase,
)
from packages.domain import AssetStatus, ClassificationStatus
router = APIRouter()
+54 -50
View File
@@ -5,36 +5,35 @@ This module is intentionally disabled until the DI container/auth ports are rebu
Do not mount it directly; use `auth_simple.py` only as the current compatibility route.
"""
raise RuntimeError(
"apps.api.app.api.routes.auth is disabled: rebuild DI container before mounting full auth routes"
)
raise RuntimeError("apps.api.app.api.routes.auth is disabled: rebuild DI container before mounting full auth routes")
from fastapi import APIRouter, HTTPException, status, Depends, Request
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel, EmailStr
from packages.application.auth import (
RegisterUserUseCase,
RegisterUserRequest,
LoginUseCase,
LoginRequest,
LogoutUseCase,
LogoutRequest,
VerifyEmailUseCase,
VerifyEmailRequest,
RequestPasswordResetUseCase,
RequestPasswordResetRequest,
ResetPasswordUseCase,
ResetPasswordRequest,
)
from packages.domain.entities import User
from apps.api.app.dependencies import get_container
from apps.api.app.middleware.auth import get_current_user
from packages.application.auth import (
LoginRequest,
LoginUseCase,
LogoutRequest,
LogoutUseCase,
RegisterUserRequest,
RegisterUserUseCase,
RequestPasswordResetRequest,
RequestPasswordResetUseCase,
ResetPasswordRequest,
ResetPasswordUseCase,
VerifyEmailRequest,
VerifyEmailUseCase,
)
from packages.domain.entities import User
router = APIRouter(prefix="/auth", tags=["Authentication"])
# ==================== Request/Response Models ====================
class RegisterRequestModel(BaseModel):
email: EmailStr
password: str
@@ -77,11 +76,16 @@ class ResetPasswordModel(BaseModel):
# ==================== API Endpoints ====================
@router.post("/register", response_model=RegisterResponseModel, status_code=status.HTTP_201_CREATED)
@router.post(
"/register",
response_model=RegisterResponseModel,
status_code=status.HTTP_201_CREATED,
)
async def register(request: RegisterRequestModel):
"""
用户注册
- 邮箱必须唯一
- 用户名必须唯一
- 密码至少 8 位,包含大小写字母和数字
@@ -89,22 +93,22 @@ async def register(request: RegisterRequestModel):
"""
container = get_container()
use_case = container.get_register_user_use_case()
req = RegisterUserRequest(
email=request.email,
password=request.password,
username=request.username,
display_name=request.display_name,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=error,
)
return RegisterResponseModel(
user_id=response.user_id,
email=response.email,
@@ -118,7 +122,7 @@ async def register(request: RegisterRequestModel):
async def login(request: LoginRequestModel):
"""
用户登录
- 使用邮箱和密码登录
- 返回 access_token 和 refresh_token
- access_token 有效期 30 分钟
@@ -126,20 +130,20 @@ async def login(request: LoginRequestModel):
"""
container = get_container()
use_case = container.get_login_use_case()
req = LoginRequest(
email=request.email,
password=request.password,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=error,
)
return LoginResponseModel(
access_token=response.access_token,
refresh_token=response.refresh_token,
@@ -159,16 +163,16 @@ async def logout(
):
"""
用户登出
- 默认只登出当前设备
- 设置 logout_all_devices=true 可登出所有设备
"""
container = get_container()
use_case = container.get_logout_use_case()
# 从 JWT token 中提取 session_id
from packages.domain.auth import jwt_service
# 从 request 中获取 token
auth_header = request.headers.get("Authorization")
session_id = None
@@ -179,15 +183,15 @@ async def logout(
session_id = payload.get("sid") # 从 payload 提取 session_id
except:
pass # token 无效或没有 session_id,继续使用 None
req = LogoutRequest(
user_id=current_user.id,
session_id=session_id,
logout_all_devices=logout_all_devices,
)
success, error = use_case.execute(req)
if not success:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@@ -199,23 +203,23 @@ async def logout(
async def verify_email(token: str):
"""
邮箱验证
- 通过邮件中的链接访问此接口
- 验证成功后标记邮箱为已验证
"""
container = get_container()
use_case = container.get_verify_email_use_case()
req = VerifyEmailRequest(token=token)
success, error = use_case.execute(req)
if not success:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=error,
)
return {"message": "Email verified successfully"}
@@ -223,18 +227,18 @@ async def verify_email(token: str):
async def forgot_password(request: PasswordResetRequestModel):
"""
请求密码重置
- 发送密码重置邮件
- 邮件中包含重置链接(有效期 1 小时)
- 即使邮箱不存在也返回成功(安全考虑)
"""
container = get_container()
use_case = container.get_request_password_reset_use_case()
req = RequestPasswordResetRequest(email=request.email)
success, error = use_case.execute(req)
# 不论成功失败都返回 202(安全考虑)
return {"message": "Password reset email sent if account exists"}
@@ -243,24 +247,24 @@ async def forgot_password(request: PasswordResetRequestModel):
async def reset_password(request: ResetPasswordModel):
"""
重置密码
- 使用邮件中的 token 重置密码
- 新密码必须符合密码强度要求
"""
container = get_container()
use_case = container.get_reset_password_use_case()
req = ResetPasswordRequest(
token=request.token,
new_password=request.new_password,
)
success, error = use_case.execute(req)
if not success:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=error,
)
return {"message": "Password reset successfully"}
+8 -4
View File
@@ -1,17 +1,18 @@
"""
认证 APISQLAlchemy ORM
"""
from datetime import datetime, timedelta, timezone
import hashlib
import secrets
from datetime import datetime, timedelta, timezone
import jwt
from app.config import settings
from app.dependencies import get_db_session
from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel, EmailStr
from sqlalchemy.orm import Session
from app.config import settings
from app.dependencies import get_db_session
from packages.adapters.sqlalchemy_impl.models import UserModel
from packages.domain.auth import password_hasher, password_validator
@@ -169,4 +170,7 @@ async def login(request: LoginRequest, db: Session = Depends(get_db_session)):
@router.get("/me")
async def get_current_user_info():
raise HTTPException(status_code=status.HTTP_501_NOT_IMPLEMENTED, detail="/auth/me requires bearer-token dependency integration")
raise HTTPException(
status_code=status.HTTP_501_NOT_IMPLEMENTED,
detail="/auth/me requires bearer-token dependency integration",
)
+11 -5
View File
@@ -1,12 +1,18 @@
from datetime import datetime, timezone
from fastapi import APIRouter, Depends, HTTPException
from typing import Any
from app.core.celery_app import celery_app
from app.dependencies import get_classification_job_repository
from app.schemas.classification_job import ClassificationJobResponse, SubmitClassificationJobRequest
from typing import Any
from packages.application import SubmitClassificationJobCommand, SubmitClassificationJobUseCase
from app.schemas.classification_job import (
ClassificationJobResponse,
SubmitClassificationJobRequest,
)
from fastapi import APIRouter, Depends, HTTPException
from packages.application import (
SubmitClassificationJobCommand,
SubmitClassificationJobUseCase,
)
router = APIRouter()
+3 -2
View File
@@ -1,4 +1,4 @@
from fastapi import APIRouter, Depends, HTTPException
from typing import Any
from app.core.storage import MinIOService, get_minio_service
from app.dependencies import get_generated_video_repository
@@ -7,7 +7,8 @@ from app.schemas.generated_video import (
GeneratedVideoResponse,
ListGeneratedVideosResponse,
)
from typing import Any
from fastapi import APIRouter, Depends, HTTPException
from packages.application import (
GetGeneratedVideoDownloadUrlUseCase,
GetGeneratedVideoUseCase,
+15 -5
View File
@@ -1,10 +1,20 @@
from fastapi import APIRouter, Depends, HTTPException
from typing import Any
from app.core.celery_app import celery_app
from app.dependencies import get_generation_task_repository, get_generated_video_repository
from app.schemas.generation_task import CreateGenerationTaskRequest, GenerationTaskResponse
from app.schemas.generated_video import GeneratedVideoResponse, ListGeneratedVideosResponse
from typing import Any
from app.dependencies import (
get_generated_video_repository,
get_generation_task_repository,
)
from app.schemas.generated_video import (
GeneratedVideoResponse,
ListGeneratedVideosResponse,
)
from app.schemas.generation_task import (
CreateGenerationTaskRequest,
GenerationTaskResponse,
)
from fastapi import APIRouter, Depends, HTTPException
from packages.application import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
+33 -13
View File
@@ -1,12 +1,11 @@
from pydantic import BaseModel
from datetime import datetime
import psycopg2
import redis
from app.config import settings
from fastapi import APIRouter, status
from fastapi.responses import JSONResponse
from app.config import settings
from pydantic import BaseModel
router = APIRouter(tags=["Health"])
@@ -56,16 +55,28 @@ async def startup_check():
async def _check_database() -> dict:
if settings.USE_IN_MEMORY_DB:
return {"status": "healthy", "type": "in_memory", "message": "Using in-memory database"}
return {
"status": "healthy",
"type": "in_memory",
"message": "Using in-memory database",
}
try:
conn = psycopg2.connect(settings.DATABASE_URL, connect_timeout=3)
with conn.cursor() as cur:
cur.execute("SELECT 1")
cur.fetchone()
conn.close()
return {"status": "healthy", "type": "postgresql", "message": "Database connection successful"}
return {
"status": "healthy",
"type": "postgresql",
"message": "Database connection successful",
}
except Exception as error:
return {"status": "unhealthy", "type": "postgresql", "message": f"Database connection failed: {error}"}
return {
"status": "unhealthy",
"type": "postgresql",
"message": f"Database connection failed: {error}",
}
async def _check_redis() -> dict:
@@ -73,23 +84,32 @@ async def _check_redis() -> dict:
client = redis.from_url(settings.REDIS_URL, socket_connect_timeout=3)
client.ping()
client.close()
return {"status": "healthy", "type": "redis", "message": "Redis connection successful"}
return {
"status": "healthy",
"type": "redis",
"message": "Redis connection successful",
}
except Exception as error:
return {"status": "unhealthy", "type": "redis", "message": f"Redis connection failed: {error}"}
return {
"status": "unhealthy",
"type": "redis",
"message": f"Redis connection failed: {error}",
}
async def _check_migrations() -> dict:
if settings.USE_IN_MEMORY_DB:
return {"status": "healthy", "message": "Using in-memory database, no migrations needed"}
return {
"status": "healthy",
"message": "Using in-memory database, no migrations needed",
}
try:
conn = psycopg2.connect(settings.DATABASE_URL, connect_timeout=3)
with conn.cursor() as cur:
cur.execute(
"""
cur.execute("""
SELECT COUNT(*) FROM information_schema.tables
WHERE table_name IN ('projects', 'asset_libraries', 'assets', 'ingest_jobs', 'classification_jobs')
"""
)
""")
count = cur.fetchone()[0]
conn.close()
if count >= 5:
+3 -2
View File
@@ -1,9 +1,10 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.core.celery_app import celery_app
from app.dependencies import get_ingest_job_repository
from app.schemas.ingest_job import IngestJobResponse, SubmitIngestJobRequest
from typing import Any
from fastapi import APIRouter, Depends
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
router = APIRouter()
@@ -1,4 +1,5 @@
"""项目管理 API 路由"""
from datetime import datetime
from typing import Annotated
@@ -11,7 +12,6 @@ from packages.adapters.sqlite_tracker.project_management_repositories import (
SQLiteTaskRepository,
)
from packages.application.get_task_detail_use_case import GetTaskDetailUseCase
from packages.application.update_task_use_case import UpdateTaskUseCase
from packages.application.project_management_use_cases import (
CreateMilestoneUseCase,
CreateTaskIssueUseCase,
@@ -23,6 +23,7 @@ from packages.application.project_management_use_cases import (
UpdateTaskProgressUseCase,
UpdateTaskStatusUseCase,
)
from packages.application.update_task_use_case import UpdateTaskUseCase
from packages.domain import TaskPriority, TaskStatus
router = APIRouter()
+13 -4
View File
@@ -1,9 +1,18 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.dependencies import get_project_repository
from app.schemas.project import CreateProjectRequest, ListProjectsResponse, ProjectResponse
from typing import Any
from packages.application import CreateProjectCommand, CreateProjectUseCase, ListProjectsUseCase
from app.schemas.project import (
CreateProjectRequest,
ListProjectsResponse,
ProjectResponse,
)
from fastapi import APIRouter, Depends
from packages.application import (
CreateProjectCommand,
CreateProjectUseCase,
ListProjectsUseCase,
)
router = APIRouter()
+3 -2
View File
@@ -1,11 +1,12 @@
from fastapi import APIRouter, Depends, File, Form, UploadFile
from typing import Any
from uuid import uuid4
from app.core.celery_app import celery_app
from app.core.storage import MinIOService, get_minio_service
from app.dependencies import get_ingest_job_repository
from app.schemas.upload import UploadAssetResponse
from typing import Any
from fastapi import APIRouter, Depends, File, Form, UploadFile
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
router = APIRouter()
+60 -49
View File
@@ -9,21 +9,28 @@ raise RuntimeError(
"apps.api.app.api.routes.workspaces is disabled: rebuild DI container before mounting workspace routes"
)
from fastapi import APIRouter, HTTPException, status, Depends
from pydantic import BaseModel, EmailStr
from typing import List
from datetime import datetime
from typing import List
from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel, EmailStr
from apps.api.app.dependencies import get_container
from apps.api.app.middleware.auth import (
get_current_user,
require_workspace_access,
require_workspace_admin,
require_workspace_owner,
)
from packages.application.workspace import *
from packages.domain.entities import User
from apps.api.app.dependencies import get_container
from apps.api.app.middleware.auth import get_current_user, require_workspace_access, require_workspace_admin, require_workspace_owner
router = APIRouter(prefix="/workspaces", tags=["Workspaces"])
# ==================== Request/Response Models ====================
class CreateWorkspaceRequestModel(BaseModel):
name: str
subscription_plan: str = "free"
@@ -52,6 +59,7 @@ class UpgradeSubscriptionRequestModel(BaseModel):
# ==================== Workspace CRUD ====================
@router.post("", response_model=WorkspaceResponseModel, status_code=status.HTTP_201_CREATED)
async def create_workspace(
request: CreateWorkspaceRequestModel,
@@ -60,18 +68,18 @@ async def create_workspace(
"""创建工作空间"""
container = get_container()
use_case = container.get_create_workspace_use_case()
req = CreateWorkspaceRequest(
name=request.name,
owner_user_id=current_user.id,
subscription_plan=request.subscription_plan,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return WorkspaceResponseModel(
workspace_id=response.workspace_id,
name=response.name,
@@ -86,13 +94,13 @@ async def list_workspaces(current_user: User = Depends(get_current_user)):
"""获取用户的所有工作空间"""
container = get_container()
use_case = container.get_list_workspaces_use_case()
req = ListWorkspacesRequest(user_id=current_user.id)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {
"workspaces": [
{
@@ -117,13 +125,13 @@ async def get_workspace_detail(
"""获取工作空间详情"""
container = get_container()
use_case = container.get_get_workspace_detail_use_case()
req = GetWorkspaceDetailRequest(workspace_id=workspace_id, user_id=current_user.id)
detail, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=error)
return {
"workspace_id": detail.workspace_id,
"name": detail.name,
@@ -140,6 +148,7 @@ async def get_workspace_detail(
# ==================== Member Management ====================
@router.post("/{workspace_id}/members/invite", status_code=status.HTTP_201_CREATED)
async def invite_member(
workspace_id: str,
@@ -149,19 +158,19 @@ async def invite_member(
"""邀请成员"""
container = get_container()
use_case = container.get_invite_member_use_case()
req = InviteMemberRequest(
workspace_id=workspace_id,
inviter_user_id=current_user.id,
invitee_email=request.email,
role=request.role,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {
"invitation_id": response.invitation_id,
"invitee_email": response.invitee_email,
@@ -178,13 +187,13 @@ async def list_members(
"""获取成员列表"""
container = get_container()
use_case = container.get_list_members_use_case()
req = ListMembersRequest(workspace_id=workspace_id, requester_user_id=current_user.id)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=error)
return {
"members": [
{
@@ -211,15 +220,15 @@ async def remove_member(
"""移除成员"""
container = get_container()
use_case = container.get_remove_member_use_case()
req = RemoveMemberRequest(
workspace_id=workspace_id,
requester_user_id=current_user.id,
target_user_id=user_id,
)
success, error = use_case.execute(req)
if not success:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
@@ -232,10 +241,10 @@ async def leave_workspace(
"""离开工作空间"""
container = get_container()
use_case = container.get_leave_workspace_use_case()
req = LeaveWorkspaceRequest(workspace_id=workspace_id, user_id=current_user.id)
success, error = use_case.execute(req)
if not success:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
@@ -250,19 +259,19 @@ async def update_member_role(
"""修改成员角色"""
container = get_container()
use_case = container.get_update_member_role_use_case()
req = UpdateMemberRoleRequest(
workspace_id=workspace_id,
requester_user_id=current_user.id,
target_user_id=user_id,
new_role=request.role,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {
"user_id": response.user_id,
"old_role": response.old_role,
@@ -272,6 +281,7 @@ async def update_member_role(
# ==================== Subscription Management ====================
@router.post("/{workspace_id}/subscription/upgrade")
async def upgrade_subscription(
workspace_id: str,
@@ -281,18 +291,18 @@ async def upgrade_subscription(
"""升级订阅"""
container = get_container()
use_case = container.get_upgrade_subscription_use_case()
req = UpgradeSubscriptionRequest(
workspace_id=workspace_id,
requester_user_id=current_user.id,
new_plan=request.new_plan,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {
"workspace_id": response.workspace_id,
"old_plan": response.old_plan,
@@ -310,13 +320,13 @@ async def cancel_subscription(
"""取消订阅"""
container = get_container()
use_case = container.get_cancel_subscription_use_case()
req = CancelSubscriptionRequest(workspace_id=workspace_id, requester_user_id=current_user.id)
success, error = use_case.execute(req)
if not success:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {"message": "Subscription cancelled successfully"}
@@ -328,24 +338,25 @@ async def get_quota_status(
"""获取配额状态"""
container = get_container()
quota_checker = container.quota_checker
# 检查权限
permission_checker = container.permission_checker
has_access, _ = permission_checker.check_workspace_access(workspace_id, current_user.id)
if not has_access:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Access denied")
status = quota_checker.get_quota_status(workspace_id)
if not status:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Workspace not found")
return status
# ==================== Invitation Acceptance ====================
@router.post("/invitations/{token}/accept")
async def accept_invitation(
token: str,
@@ -354,13 +365,13 @@ async def accept_invitation(
"""接受邀请"""
container = get_container()
use_case = container.get_accept_invitation_use_case()
req = AcceptInvitationRequest(invitation_token=token, user_id=current_user.id)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {
"workspace_id": response.workspace_id,
"workspace_name": response.workspace_name,
@@ -373,11 +384,11 @@ async def decline_invitation(token: str):
"""拒绝邀请"""
container = get_container()
use_case = container.get_decline_invitation_use_case()
req = DeclineInvitationRequest(invitation_token=token)
success, error = use_case.execute(req)
if not success:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {"message": "Invitation declined"}
+3 -2
View File
@@ -1,6 +1,7 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
from typing import Optional
import os
from typing import Optional
from pydantic_settings import BaseSettings, SettingsConfigDict
class Settings(BaseSettings):
+1 -3
View File
@@ -1,7 +1,5 @@
from celery import Celery
from app.config import get_settings
from celery import Celery
settings = get_settings()
celery_app = Celery("xiaoxia-saas-api")
+3 -2
View File
@@ -1,6 +1,7 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
from typing import Optional
import os
from typing import Optional
from pydantic_settings import BaseSettings, SettingsConfigDict
class DatabaseSettings(BaseSettings):
+18 -22
View File
@@ -1,7 +1,8 @@
"""阿里云 OSS 存储服务"""
import logging
from urllib.parse import urlparse
import os
from urllib.parse import urlparse
try:
import oss2
@@ -48,12 +49,12 @@ class OSSStorageService:
) -> str:
"""
上传文件到 OSS
Args:
file_or_path: 文件对象或本地文件路径
storage_key: 存储键(文件路径)
content_type: 内容类型
Returns:
文件公网 URL
"""
@@ -63,20 +64,12 @@ class OSSStorageService:
try:
# 如果是字符串路径,从本地文件上传
if isinstance(file_or_path, str):
self.bucket.put_object_from_file(
storage_key,
file_or_path,
headers={'Content-Type': content_type}
)
self.bucket.put_object_from_file(storage_key, file_or_path, headers={"Content-Type": content_type})
else:
# 文件对象
file_or_path.seek(0)
self.bucket.put_object(
storage_key,
file_or_path,
headers={'Content-Type': content_type}
)
self.bucket.put_object(storage_key, file_or_path, headers={"Content-Type": content_type})
return f"{self.public_url}/{storage_key}"
except Exception as e:
raise Exception(f"Failed to upload file to OSS: {e}")
@@ -88,11 +81,11 @@ class OSSStorageService:
def get_download_url(self, storage_key_or_url: str, expires_seconds: int = 3600) -> str:
"""
获取文件下载签名 URL(用于私有文件)
Args:
storage_key_or_url: 存储键或完整 URL
expires_seconds: 过期时间(秒)
Returns:
签名 URL
"""
@@ -103,7 +96,7 @@ class OSSStorageService:
storage_key = self._normalize_storage_key(storage_key_or_url)
try:
return self.bucket.sign_url('GET', storage_key, expires_seconds)
return self.bucket.sign_url("GET", storage_key, expires_seconds)
except Exception:
return self.get_url(storage_key)
@@ -118,7 +111,7 @@ class OSSStorageService:
def download_file(self, storage_key: str, local_path: str):
"""
从 OSS 下载文件到本地
Args:
storage_key: 存储键
local_path: 本地文件路径
@@ -135,7 +128,7 @@ class OSSStorageService:
def delete_file(self, storage_key: str):
"""
删除 OSS 文件
Args:
storage_key: 存储键
"""
@@ -145,15 +138,18 @@ class OSSStorageService:
try:
self.bucket.delete_object(storage_key)
except Exception as error:
logger.warning("Failed to delete file from OSS", extra={"storage_key": storage_key, "error": str(error)})
logger.warning(
"Failed to delete file from OSS",
extra={"storage_key": storage_key, "error": str(error)},
)
def file_exists(self, storage_key: str) -> bool:
"""
检查文件是否存在
Args:
storage_key: 存储键
Returns:
是否存在
"""
+6 -2
View File
@@ -1,9 +1,13 @@
from collections.abc import Generator
from app.config import settings
from sqlalchemy.orm import Session
from app.config import settings
from packages.adapters.sqlalchemy_impl import build_session_factory, ensure_database_exists, initialize_database
from packages.adapters.sqlalchemy_impl import (
build_session_factory,
ensure_database_exists,
initialize_database,
)
ensure_database_exists(settings.DATABASE_URL)
engine, SessionLocal = build_session_factory(
+40 -14
View File
@@ -1,14 +1,26 @@
from app.config import settings
from fastapi import Depends
from sqlalchemy.orm import Session
from app.config import settings
from packages.adapters.sqlalchemy_impl.asset_library_repository import SQLAlchemyAssetLibraryRepository
from packages.adapters.sqlalchemy_impl.asset_library_repository import (
SQLAlchemyAssetLibraryRepository,
)
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from packages.adapters.sqlalchemy_impl.classification_job_repository import SQLAlchemyClassificationJobRepository
from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository
from packages.adapters.sqlalchemy_impl.generation_task_repository import SQLAlchemyGenerationTaskRepository
from packages.adapters.sqlalchemy_impl.ingest_job_repository import SQLAlchemyIngestJobRepository
from packages.adapters.sqlalchemy_impl.project_repository import SQLAlchemyProjectRepository
from packages.adapters.sqlalchemy_impl.classification_job_repository import (
SQLAlchemyClassificationJobRepository,
)
from packages.adapters.sqlalchemy_impl.generated_video_repository import (
SQLAlchemyGeneratedVideoRepository,
)
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
from packages.adapters.sqlalchemy_impl.ingest_job_repository import (
SQLAlchemyIngestJobRepository,
)
from packages.adapters.sqlalchemy_impl.project_repository import (
SQLAlchemyProjectRepository,
)
from packages.adapters.sqlalchemy_impl.session import build_session_factory
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
@@ -22,29 +34,43 @@ def get_db_session():
session.close()
def get_asset_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyAssetRepository:
def get_asset_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyAssetRepository:
return SQLAlchemyAssetRepository(session)
def get_asset_library_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyAssetLibraryRepository:
def get_asset_library_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyAssetLibraryRepository:
return SQLAlchemyAssetLibraryRepository(session)
def get_ingest_job_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyIngestJobRepository:
def get_ingest_job_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyIngestJobRepository:
return SQLAlchemyIngestJobRepository(session)
def get_classification_job_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyClassificationJobRepository:
def get_classification_job_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyClassificationJobRepository:
return SQLAlchemyClassificationJobRepository(session)
def get_generation_task_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyGenerationTaskRepository:
def get_generation_task_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyGenerationTaskRepository:
return SQLAlchemyGenerationTaskRepository(session)
def get_generated_video_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyGeneratedVideoRepository:
def get_generated_video_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyGeneratedVideoRepository:
return SQLAlchemyGeneratedVideoRepository(session)
def get_project_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyProjectRepository:
def get_project_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyProjectRepository:
return SQLAlchemyProjectRepository(session)
+31 -34
View File
@@ -5,17 +5,14 @@ Disabled because it depends on the removed DI container. Rebuild it around the
canonical JWT settings and SQLAlchemy-backed user repository before reuse.
"""
raise RuntimeError(
"apps.api.app.middleware.auth is disabled: rebuild auth dependency wiring before importing it"
)
raise RuntimeError("apps.api.app.middleware.auth is disabled: rebuild auth dependency wiring before importing it")
from fastapi import Depends, HTTPException, status
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from apps.api.app.dependencies import get_container
from packages.domain.auth import jwt_service
from packages.domain.entities import User
from apps.api.app.dependencies import get_container
security = HTTPBearer()
@@ -25,42 +22,42 @@ async def get_current_user(
) -> User:
"""
获取当前登录用户
从 Authorization header 中提取 JWT token 并验证
Raises:
HTTPException: Token 无效或过期
Returns:
当前用户对象
"""
token = credentials.credentials
try:
# 验证 token
payload = jwt_service.verify_token(token)
user_id = payload.get("sub")
if not user_id:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid token: missing user_id",
headers={"WWW-Authenticate": "Bearer"},
)
# 从数据库获取用户
container = get_container()
user = container.user_repository.find_by_id(user_id)
if not user:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="User not found",
headers={"WWW-Authenticate": "Bearer"},
)
return user
except Exception as e:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
@@ -74,15 +71,15 @@ async def get_current_user_optional(
) -> User | None:
"""
获取当前登录用户(可选)
如果没有提供 token,返回 None 而不是抛出异常
Returns:
当前用户对象或 None
"""
if not credentials:
return None
try:
return await get_current_user(credentials)
except HTTPException:
@@ -92,78 +89,78 @@ async def get_current_user_optional(
def require_workspace_access(workspace_id: str, user: User = Depends(get_current_user)) -> tuple[str, str]:
"""
要求用户可以访问指定工作空间
Args:
workspace_id: 工作空间 ID
user: 当前用户
Raises:
HTTPException: 用户没有访问权限
Returns:
(workspace_id, user_role)
"""
container = get_container()
permission_checker = container.permission_checker
has_access, role = permission_checker.check_workspace_access(workspace_id, user.id)
if not has_access:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="You don't have access to this workspace",
)
return workspace_id, role
def require_workspace_admin(workspace_id: str, user: User = Depends(get_current_user)) -> str:
"""
要求用户是工作空间的 Admin 或 Owner
Args:
workspace_id: 工作空间 ID
user: 当前用户
Raises:
HTTPException: 用户没有管理权限
Returns:
workspace_id
"""
container = get_container()
permission_checker = container.permission_checker
if not permission_checker.check_is_admin_or_owner(workspace_id, user.id):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only workspace owner or admin can perform this action",
)
return workspace_id
def require_workspace_owner(workspace_id: str, user: User = Depends(get_current_user)) -> str:
"""
要求用户是工作空间的 Owner
Args:
workspace_id: 工作空间 ID
user: 当前用户
Raises:
HTTPException: 用户不是 Owner
Returns:
workspace_id
"""
container = get_container()
permission_checker = container.permission_checker
if not permission_checker.check_is_owner(workspace_id, user.id):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only workspace owner can perform this action",
)
return workspace_id
+21 -17
View File
@@ -1,19 +1,21 @@
"""
全局异常处理和错误响应
"""
from fastapi import Request, status
from fastapi.responses import JSONResponse
from fastapi.exceptions import RequestValidationError
from starlette.exceptions import HTTPException as StarletteHTTPException
import traceback
import logging
import traceback
from fastapi import Request, status
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse
from starlette.exceptions import HTTPException as StarletteHTTPException
logger = logging.getLogger(__name__)
class APIException(Exception):
"""API 异常基类"""
def __init__(
self,
message: str,
@@ -28,7 +30,7 @@ class APIException(Exception):
class AuthenticationError(APIException):
"""认证错误"""
def __init__(self, message: str = "Authentication failed"):
super().__init__(
message=message,
@@ -39,7 +41,7 @@ class AuthenticationError(APIException):
class PermissionDeniedError(APIException):
"""权限拒绝"""
def __init__(self, message: str = "Permission denied"):
super().__init__(
message=message,
@@ -50,7 +52,7 @@ class PermissionDeniedError(APIException):
class ResourceNotFoundError(APIException):
"""资源不存在"""
def __init__(self, resource: str = "Resource"):
super().__init__(
message=f"{resource} not found",
@@ -61,7 +63,7 @@ class ResourceNotFoundError(APIException):
class ValidationError(APIException):
"""验证错误"""
def __init__(self, message: str):
super().__init__(
message=message,
@@ -100,12 +102,14 @@ async def validation_exception_handler(request: Request, exc: RequestValidationE
"""请求验证异常处理"""
errors = []
for error in exc.errors():
errors.append({
"field": ".".join(str(loc) for loc in error["loc"]),
"message": error["msg"],
"type": error["type"],
})
errors.append(
{
"field": ".".join(str(loc) for loc in error["loc"]),
"message": error["msg"],
"type": error["type"],
}
)
return JSONResponse(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
content={
@@ -121,7 +125,7 @@ async def validation_exception_handler(request: Request, exc: RequestValidationE
async def general_exception_handler(request: Request, exc: Exception):
"""通用异常处理"""
logger.error(f"Unhandled exception: {exc}", exc_info=True)
# 生产环境不返回详细错误信息
return JSONResponse(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
+23 -24
View File
@@ -1,8 +1,10 @@
"""
请求日志中间件
"""
import time
import logging
import time
from fastapi import Request
from starlette.middleware.base import BaseHTTPMiddleware
@@ -11,58 +13,57 @@ logger = logging.getLogger(__name__)
class RequestLoggingMiddleware(BaseHTTPMiddleware):
"""请求日志中间件"""
async def dispatch(self, request: Request, call_next):
# 记录请求开始时间
start_time = time.time()
# 记录请求信息
logger.info(f"Request: {request.method} {request.url.path}")
# 处理请求
response = await call_next(request)
# 计算处理时间
process_time = time.time() - start_time
# 记录响应信息
logger.info(
f"Response: {request.method} {request.url.path} "
f"status={response.status_code} time={process_time:.3f}s"
f"Response: {request.method} {request.url.path} " f"status={response.status_code} time={process_time:.3f}s"
)
# 添加响应头
response.headers["X-Process-Time"] = str(process_time)
return response
class RateLimitMiddleware(BaseHTTPMiddleware):
"""简单的速率限制中间件(基于内存)"""
def __init__(self, app, max_requests: int = 100, window_seconds: int = 60):
super().__init__(app)
self.max_requests = max_requests
self.window_seconds = window_seconds
self.requests = {} # {ip: [(timestamp, ...)]}
async def dispatch(self, request: Request, call_next):
# 获取客户端 IP
client_ip = request.client.host
current_time = time.time()
# 清理过期记录
if client_ip in self.requests:
self.requests[client_ip] = [
ts for ts in self.requests[client_ip]
if current_time - ts < self.window_seconds
ts for ts in self.requests[client_ip] if current_time - ts < self.window_seconds
]
# 检查速率限制
request_count = len(self.requests.get(client_ip, []))
if request_count >= self.max_requests:
from fastapi.responses import JSONResponse
return JSONResponse(
status_code=429,
content={
@@ -72,19 +73,17 @@ class RateLimitMiddleware(BaseHTTPMiddleware):
}
},
)
# 记录请求
if client_ip not in self.requests:
self.requests[client_ip] = []
self.requests[client_ip].append(current_time)
# 处理请求
response = await call_next(request)
# 添加速率限制信息到响应头
response.headers["X-RateLimit-Limit"] = str(self.max_requests)
response.headers["X-RateLimit-Remaining"] = str(
self.max_requests - len(self.requests[client_ip])
)
response.headers["X-RateLimit-Remaining"] = str(self.max_requests - len(self.requests[client_ip]))
return response
+28 -26
View File
@@ -1,9 +1,11 @@
"""
性能监控中间件
"""
import time
import logging
import time
from typing import Callable
from fastapi import Request, Response
from starlette.middleware.base import BaseHTTPMiddleware
@@ -12,30 +14,30 @@ logger = logging.getLogger(__name__)
class PerformanceMonitoringMiddleware(BaseHTTPMiddleware):
"""性能监控中间件"""
def __init__(self, app, slow_request_threshold: float = 1.0):
super().__init__(app)
self.slow_request_threshold = slow_request_threshold # 慢请求阈值(秒)
async def dispatch(self, request: Request, call_next: Callable):
# 记录请求开始时间
start_time = time.time()
# 生成请求 ID
request_id = self._generate_request_id()
request.state.request_id = request_id
# 处理请求
try:
response = await call_next(request)
# 计算处理时间
process_time = time.time() - start_time
# 添加响应头
response.headers["X-Request-ID"] = request_id
response.headers["X-Process-Time"] = f"{process_time:.3f}"
# 记录慢请求
if process_time > self.slow_request_threshold:
logger.warning(
@@ -43,55 +45,55 @@ class PerformanceMonitoringMiddleware(BaseHTTPMiddleware):
f"took {process_time:.3f}s (threshold: {self.slow_request_threshold}s) "
f"[request_id={request_id}]"
)
# 记录请求日志
logger.info(
f"{request.method} {request.url.path} "
f"status={response.status_code} time={process_time:.3f}s "
f"[request_id={request_id}]"
)
return response
except Exception as e:
process_time = time.time() - start_time
logger.error(
f"Request failed: {request.method} {request.url.path} "
f"error={str(e)} time={process_time:.3f}s "
f"[request_id={request_id}]",
exc_info=True
exc_info=True,
)
raise
def _generate_request_id(self) -> str:
"""生成请求 ID"""
import uuid
return str(uuid.uuid4())
class DatabaseQueryLogger:
"""数据库查询日志记录器"""
def __init__(self):
self.queries = []
self.total_time = 0
def log_query(self, query: str, params: tuple, duration: float):
"""记录查询"""
self.queries.append({
"query": query,
"params": params,
"duration": duration,
})
self.queries.append(
{
"query": query,
"params": params,
"duration": duration,
}
)
self.total_time += duration
# 记录慢查询(超过 100ms
if duration > 0.1:
logger.warning(
f"Slow query detected: {query[:100]}... "
f"took {duration:.3f}s with params {params}"
)
logger.warning(f"Slow query detected: {query[:100]}... " f"took {duration:.3f}s with params {params}")
def get_stats(self):
"""获取统计信息"""
return {
+21 -20
View File
@@ -1,14 +1,16 @@
"""
API 版本管理中间件
"""
from datetime import datetime
from fastapi import Request
from starlette.middleware.base import BaseHTTPMiddleware
from datetime import datetime
class APIVersionMiddleware(BaseHTTPMiddleware):
"""API 版本管理中间件"""
# 版本配置
VERSIONS = {
"v1": {
@@ -24,33 +26,31 @@ class APIVersionMiddleware(BaseHTTPMiddleware):
"release_date": None,
},
}
async def dispatch(self, request: Request, call_next):
# 提取版本号
version = self._extract_version(request.url.path)
# 处理请求
response = await call_next(request)
# 添加版本信息头
if version:
response.headers["X-API-Version"] = version
# 添加弃用警告
version_info = self.VERSIONS.get(version, {})
if version_info.get("deprecated"):
response.headers["X-API-Deprecated"] = "true"
sunset_date = version_info.get("sunset_date")
if sunset_date:
response.headers["X-API-Sunset-Date"] = sunset_date
response.headers["X-API-Deprecation-Info"] = (
f"https://docs.xiaoxia-saas.com/api/deprecation/{version}"
)
response.headers["X-API-Deprecation-Info"] = f"https://docs.xiaoxia-saas.com/api/deprecation/{version}"
return response
def _extract_version(self, path: str) -> str:
"""从路径中提取版本号"""
parts = path.split("/")
@@ -62,14 +62,15 @@ class APIVersionMiddleware(BaseHTTPMiddleware):
class VersionNotFoundMiddleware(BaseHTTPMiddleware):
"""处理已下线的 API 版本"""
SUNSET_VERSIONS = [] # 已下线的版本列表
async def dispatch(self, request: Request, call_next):
version = self._extract_version(request.url.path)
if version in self.SUNSET_VERSIONS:
from fastapi.responses import JSONResponse
return JSONResponse(
status_code=410,
content={
@@ -77,13 +78,13 @@ class VersionNotFoundMiddleware(BaseHTTPMiddleware):
"code": "API_VERSION_SUNSET",
"message": f"API {version} has been sunset and is no longer available",
"sunset_date": "2028-07-01",
"migration_guide": f"https://docs.xiaoxia-saas.com/api/migration/{version}"
"migration_guide": f"https://docs.xiaoxia-saas.com/api/migration/{version}",
}
}
},
)
return await call_next(request)
def _extract_version(self, path: str) -> str:
"""从路径中提取版本号"""
parts = path.split("/")
+5 -1
View File
@@ -1,7 +1,11 @@
"""Schema package."""
from .asset import AssetResponse, CreateAssetRequest, ListAssetsResponse
from .asset_library import AssetLibraryResponse, CreateAssetLibraryRequest, ListAssetLibrariesResponse
from .asset_library import (
AssetLibraryResponse,
CreateAssetLibraryRequest,
ListAssetLibrariesResponse,
)
from .health import HealthResponse
from .ingest_job import IngestJobResponse, SubmitIngestJobRequest
from .project import CreateProjectRequest, ListProjectsResponse, ProjectResponse
+6 -7
View File
@@ -1,12 +1,5 @@
import os
from fastapi import FastAPI
from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from fastapi.middleware.gzip import GZipMiddleware
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.staticfiles import StaticFiles
from app.api.router import api_router, health_router
from app.config import settings
from app.middleware.exceptions import (
@@ -17,6 +10,12 @@ from app.middleware.exceptions import (
validation_exception_handler,
)
from app.middleware.logging import RequestLoggingMiddleware
from fastapi import FastAPI
from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from fastapi.middleware.gzip import GZipMiddleware
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.staticfiles import StaticFiles
app = FastAPI(
title="小虾 SaaS API",