style: normalize python formatting gates
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"}
|
||||
|
||||
@@ -1,17 +1,18 @@
|
||||
"""
|
||||
认证 API(SQLAlchemy 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",
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"}
|
||||
|
||||
@@ -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,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")
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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("/")
|
||||
|
||||
@@ -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
@@ -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",
|
||||
|
||||
+28
-18
@@ -4,13 +4,22 @@ from datetime import datetime, timezone
|
||||
|
||||
from app.config import get_settings
|
||||
from app.core.storage import get_minio_service
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
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.session import (
|
||||
SessionLocal,
|
||||
build_session_factory,
|
||||
)
|
||||
from packages.domain import GeneratedVideo, GenerationTaskStatus
|
||||
|
||||
from .celery_app import celery_app
|
||||
from .video_processing import VideoProcessor
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal, build_session_factory
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import SQLAlchemyGenerationTaskRepository
|
||||
from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.domain import GeneratedVideo, GenerationTaskStatus
|
||||
|
||||
settings = get_settings()
|
||||
if SessionLocal is None:
|
||||
@@ -21,7 +30,7 @@ if SessionLocal is None:
|
||||
def generate_video(task_id: str) -> dict:
|
||||
session = SessionLocal()
|
||||
temp_dir = tempfile.mkdtemp()
|
||||
|
||||
|
||||
try:
|
||||
task_repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
video_repo = SQLAlchemyGeneratedVideoRepository(session)
|
||||
@@ -44,7 +53,7 @@ def generate_video(task_id: str) -> dict:
|
||||
assets = asset_repo.list_by_library(task.asset_library_id)
|
||||
if not assets:
|
||||
raise RuntimeError(f"No assets found in library {task.asset_library_id}")
|
||||
|
||||
|
||||
task.progress = 20.0
|
||||
task_repo.update(task)
|
||||
session.commit()
|
||||
@@ -53,13 +62,13 @@ def generate_video(task_id: str) -> dict:
|
||||
video_assets = [a for a in assets if a.mime_type.startswith("video/")][:3]
|
||||
if not video_assets:
|
||||
raise RuntimeError("No video assets found")
|
||||
|
||||
|
||||
local_paths = []
|
||||
for i, asset in enumerate(video_assets):
|
||||
local_path = os.path.join(temp_dir, f"input_{i}.mp4")
|
||||
storage_service.download_file(asset.storage_key, local_path)
|
||||
local_paths.append(local_path)
|
||||
|
||||
|
||||
task.progress = 20.0 + (i + 1) * 10.0
|
||||
task_repo.update(task)
|
||||
session.commit()
|
||||
@@ -68,18 +77,18 @@ def generate_video(task_id: str) -> dict:
|
||||
processor = VideoProcessor(temp_dir=temp_dir)
|
||||
output_filename = f"{task.id}.mp4"
|
||||
output_path = os.path.join(temp_dir, output_filename)
|
||||
|
||||
|
||||
task.progress = 50.0
|
||||
task_repo.update(task)
|
||||
session.commit()
|
||||
|
||||
|
||||
result = processor.concatenate_videos(
|
||||
input_paths=local_paths,
|
||||
output_path=output_path,
|
||||
resolution=(1920, 1080),
|
||||
fps=25,
|
||||
)
|
||||
|
||||
|
||||
task.progress = 80.0
|
||||
task_repo.update(task)
|
||||
session.commit()
|
||||
@@ -87,13 +96,13 @@ def generate_video(task_id: str) -> dict:
|
||||
# 6. 上传到 MinIO
|
||||
storage_key = f"workspaces/{task.workspace_id}/projects/{task.project_id}/generated/{task.id}/{output_filename}"
|
||||
thumbnail_key = f"workspaces/{task.workspace_id}/projects/{task.project_id}/generated/{task.id}/thumbnail.jpg"
|
||||
|
||||
|
||||
storage_service.upload_file(result.output_path, storage_key)
|
||||
storage_service.upload_file(result.thumbnail_path, thumbnail_key)
|
||||
|
||||
|
||||
file_url = storage_service.get_url(storage_key)
|
||||
thumbnail_url = storage_service.get_url(thumbnail_key)
|
||||
|
||||
|
||||
task.progress = 90.0
|
||||
task_repo.update(task)
|
||||
session.commit()
|
||||
@@ -130,7 +139,7 @@ def generate_video(task_id: str) -> dict:
|
||||
"duration": result.duration,
|
||||
"file_size": result.file_size,
|
||||
}
|
||||
|
||||
|
||||
except Exception as error:
|
||||
try:
|
||||
task_repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
@@ -143,14 +152,15 @@ def generate_video(task_id: str) -> dict:
|
||||
session.commit()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
return {"ok": False, "task_id": task_id, "error": str(error)}
|
||||
|
||||
|
||||
finally:
|
||||
session.close()
|
||||
# 清理临时文件
|
||||
try:
|
||||
import shutil
|
||||
|
||||
shutil.rmtree(temp_dir, ignore_errors=True)
|
||||
except:
|
||||
pass
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""
|
||||
视频处理模块
|
||||
"""
|
||||
|
||||
from .processor import VideoProcessor, VideoResult
|
||||
|
||||
__all__ = ["VideoProcessor", "VideoResult"]
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
"""
|
||||
视频处理核心类
|
||||
"""
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
@@ -13,6 +14,7 @@ import ffmpeg
|
||||
@dataclass
|
||||
class VideoResult:
|
||||
"""视频生成结果"""
|
||||
|
||||
output_path: str
|
||||
thumbnail_path: str
|
||||
duration: float
|
||||
@@ -24,16 +26,16 @@ class VideoResult:
|
||||
|
||||
class VideoProcessor:
|
||||
"""视频处理器"""
|
||||
|
||||
|
||||
def __init__(self, temp_dir: str = None):
|
||||
"""
|
||||
初始化视频处理器
|
||||
|
||||
|
||||
Args:
|
||||
temp_dir: 临时文件目录,默认使用系统临时目录
|
||||
"""
|
||||
self.temp_dir = temp_dir or tempfile.gettempdir()
|
||||
|
||||
|
||||
def concatenate_videos(
|
||||
self,
|
||||
input_paths: List[str],
|
||||
@@ -43,22 +45,22 @@ class VideoProcessor:
|
||||
) -> VideoResult:
|
||||
"""
|
||||
拼接多个视频
|
||||
|
||||
|
||||
Args:
|
||||
input_paths: 输入视频路径列表
|
||||
output_path: 输出视频路径
|
||||
resolution: 输出分辨率 (width, height)
|
||||
fps: 输出帧率
|
||||
|
||||
|
||||
Returns:
|
||||
VideoResult: 生成结果
|
||||
"""
|
||||
if not input_paths:
|
||||
raise ValueError("input_paths cannot be empty")
|
||||
|
||||
|
||||
# 确保输出目录存在
|
||||
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
||||
|
||||
|
||||
try:
|
||||
# 创建临时文件列表
|
||||
concat_file = os.path.join(self.temp_dir, f"concat_{os.getpid()}.txt")
|
||||
@@ -66,12 +68,11 @@ class VideoProcessor:
|
||||
for path in input_paths:
|
||||
# FFmpeg concat demuxer 格式
|
||||
f.write(f"file '{os.path.abspath(path)}'\n")
|
||||
|
||||
|
||||
# 使用 FFmpeg 拼接视频
|
||||
width, height = resolution
|
||||
(
|
||||
ffmpeg
|
||||
.input(concat_file, format="concat", safe=0)
|
||||
ffmpeg.input(concat_file, format="concat", safe=0)
|
||||
.output(
|
||||
output_path,
|
||||
vcodec="libx264",
|
||||
@@ -84,28 +85,28 @@ class VideoProcessor:
|
||||
.overwrite_output()
|
||||
.run(capture_stdout=True, capture_stderr=True)
|
||||
)
|
||||
|
||||
|
||||
# 清理临时文件
|
||||
os.remove(concat_file)
|
||||
|
||||
|
||||
# 获取视频元数据
|
||||
probe = ffmpeg.probe(output_path)
|
||||
video_info = next(s for s in probe["streams"] if s["codec_type"] == "video")
|
||||
|
||||
|
||||
duration = float(probe["format"]["duration"])
|
||||
width = int(video_info["width"])
|
||||
height = int(video_info["height"])
|
||||
|
||||
|
||||
# 计算帧率
|
||||
fps_str = video_info.get("r_frame_rate", "25/1")
|
||||
fps_parts = fps_str.split("/")
|
||||
fps_value = float(fps_parts[0]) / float(fps_parts[1]) if len(fps_parts) == 2 else float(fps_parts[0])
|
||||
|
||||
|
||||
file_size = os.path.getsize(output_path)
|
||||
|
||||
|
||||
# 生成缩略图
|
||||
thumbnail_path = self.generate_thumbnail(output_path)
|
||||
|
||||
|
||||
return VideoResult(
|
||||
output_path=output_path,
|
||||
thumbnail_path=thumbnail_path,
|
||||
@@ -115,11 +116,11 @@ class VideoProcessor:
|
||||
fps=fps_value,
|
||||
file_size=file_size,
|
||||
)
|
||||
|
||||
|
||||
except ffmpeg.Error as e:
|
||||
stderr = e.stderr.decode() if e.stderr else ""
|
||||
raise RuntimeError(f"FFmpeg error: {stderr}") from e
|
||||
|
||||
|
||||
def generate_thumbnail(
|
||||
self,
|
||||
video_path: str,
|
||||
@@ -128,55 +129,54 @@ class VideoProcessor:
|
||||
) -> str:
|
||||
"""
|
||||
生成视频缩略图
|
||||
|
||||
|
||||
Args:
|
||||
video_path: 视频文件路径
|
||||
timestamp: 截图时间点(秒)
|
||||
output_path: 输出路径,默认为视频路径 + .jpg
|
||||
|
||||
|
||||
Returns:
|
||||
缩略图路径
|
||||
"""
|
||||
if output_path is None:
|
||||
output_path = f"{os.path.splitext(video_path)[0]}_thumb.jpg"
|
||||
|
||||
|
||||
try:
|
||||
(
|
||||
ffmpeg
|
||||
.input(video_path, ss=timestamp)
|
||||
ffmpeg.input(video_path, ss=timestamp)
|
||||
.output(output_path, vframes=1, format="image2", vcodec="mjpeg")
|
||||
.overwrite_output()
|
||||
.run(capture_stdout=True, capture_stderr=True)
|
||||
)
|
||||
|
||||
|
||||
return output_path
|
||||
|
||||
|
||||
except ffmpeg.Error as e:
|
||||
stderr = e.stderr.decode() if e.stderr else ""
|
||||
raise RuntimeError(f"FFmpeg thumbnail error: {stderr}") from e
|
||||
|
||||
|
||||
def get_video_info(self, video_path: str) -> dict:
|
||||
"""
|
||||
获取视频信息
|
||||
|
||||
|
||||
Args:
|
||||
video_path: 视频文件路径
|
||||
|
||||
|
||||
Returns:
|
||||
视频元数据字典
|
||||
"""
|
||||
try:
|
||||
probe = ffmpeg.probe(video_path)
|
||||
video_info = next(s for s in probe["streams"] if s["codec_type"] == "video")
|
||||
|
||||
|
||||
duration = float(probe["format"]["duration"])
|
||||
width = int(video_info["width"])
|
||||
height = int(video_info["height"])
|
||||
|
||||
|
||||
fps_str = video_info.get("r_frame_rate", "25/1")
|
||||
fps_parts = fps_str.split("/")
|
||||
fps_value = float(fps_parts[0]) / float(fps_parts[1]) if len(fps_parts) == 2 else float(fps_parts[0])
|
||||
|
||||
|
||||
return {
|
||||
"duration": duration,
|
||||
"width": width,
|
||||
@@ -185,7 +185,7 @@ class VideoProcessor:
|
||||
"codec": video_info.get("codec_name"),
|
||||
"bitrate": int(probe["format"].get("bit_rate", 0)),
|
||||
}
|
||||
|
||||
|
||||
except ffmpeg.Error as e:
|
||||
stderr = e.stderr.decode() if e.stderr else ""
|
||||
raise RuntimeError(f"FFmpeg probe error: {stderr}") from e
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
from celery import Celery
|
||||
|
||||
from worker_app.core.config import get_settings
|
||||
|
||||
|
||||
settings = get_settings()
|
||||
celery_app = Celery(settings.worker_name)
|
||||
celery_app.conf.broker_url = settings.broker_url
|
||||
|
||||
@@ -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 WorkerSettings(BaseSettings):
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
from worker_app.core.config import get_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,
|
||||
)
|
||||
|
||||
settings = get_settings()
|
||||
ensure_database_exists(settings.database_url)
|
||||
|
||||
@@ -1,14 +1,21 @@
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
from packages.domain import AssetClassification, ClassificationJob, ClassificationJobStatus
|
||||
from packages.adapters.sqlalchemy_impl.classification_job_repository import SQLAlchemyClassificationJobRepository
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.classification_job_repository import (
|
||||
SQLAlchemyClassificationJobRepository,
|
||||
)
|
||||
from packages.domain import (
|
||||
AssetClassification,
|
||||
ClassificationJob,
|
||||
ClassificationJobStatus,
|
||||
)
|
||||
|
||||
|
||||
@celery_app.task(name="worker.classify_asset")
|
||||
def classify_asset(job_id: str) -> dict:
|
||||
"""
|
||||
Classify asset task.
|
||||
|
||||
|
||||
Steps:
|
||||
1. Fetch ClassificationJob from repository
|
||||
2. Fetch Asset from repository
|
||||
@@ -20,31 +27,31 @@ def classify_asset(job_id: str) -> dict:
|
||||
session = SessionLocal()
|
||||
try:
|
||||
job_repo = SQLAlchemyClassificationJobRepository(session)
|
||||
|
||||
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
return {"status": "failed", "error": "job not found"}
|
||||
|
||||
|
||||
try:
|
||||
# Update job status to PROCESSING
|
||||
job.status = ClassificationJobStatus.PROCESSING
|
||||
job_repo.update(job)
|
||||
session.commit()
|
||||
|
||||
|
||||
# Mock classification (in real implementation: use ML model, vision API, etc.)
|
||||
# For now, randomly classify based on asset_id hash
|
||||
asset_id_hash = sum(ord(c) for c in job.asset_id)
|
||||
classifications = list(AssetClassification)
|
||||
classification = classifications[asset_id_hash % len(classifications)]
|
||||
confidence = 0.85
|
||||
|
||||
|
||||
# Update job status to COMPLETED
|
||||
job.status = ClassificationJobStatus.COMPLETED
|
||||
job.classification = classification.value
|
||||
job.confidence = confidence
|
||||
job_repo.update(job)
|
||||
session.commit()
|
||||
|
||||
|
||||
return {
|
||||
"status": "completed",
|
||||
"job_id": job.id,
|
||||
@@ -58,7 +65,7 @@ def classify_asset(job_id: str) -> dict:
|
||||
job.error_message = str(e)
|
||||
job_repo.update(job)
|
||||
session.commit()
|
||||
|
||||
|
||||
return {
|
||||
"status": "failed",
|
||||
"job_id": job.id,
|
||||
|
||||
@@ -7,6 +7,8 @@ from pathlib import Path
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import oss2
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl import (
|
||||
SQLAlchemyAssetRepository,
|
||||
@@ -14,8 +16,6 @@ from packages.adapters.sqlalchemy_impl import (
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.domain import GeneratedVideo, GenerationTaskStatus
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
OUTPUT_WIDTH = 1280
|
||||
OUTPUT_HEIGHT = 720
|
||||
@@ -176,7 +176,11 @@ def generate_video(task_id: str) -> dict:
|
||||
task = task_repo.get(task_id)
|
||||
if task is None:
|
||||
db.close()
|
||||
return {"status": "failed", "error": "generation task not found", "task_id": task_id}
|
||||
return {
|
||||
"status": "failed",
|
||||
"error": "generation task not found",
|
||||
"task_id": task_id,
|
||||
}
|
||||
|
||||
try:
|
||||
task.status = GenerationTaskStatus.RUNNING
|
||||
@@ -184,10 +188,14 @@ def generate_video(task_id: str) -> dict:
|
||||
task.started_at = task.started_at or datetime.now(timezone.utc)
|
||||
task_repo.update(task)
|
||||
|
||||
assets = [asset for asset in asset_repo.list_by_library(task.asset_library_id) if asset.mime_type.startswith("video")]
|
||||
assets = [
|
||||
asset for asset in asset_repo.list_by_library(task.asset_library_id) if asset.mime_type.startswith("video")
|
||||
]
|
||||
|
||||
output_name = f"generated-{task.id}.mp4"
|
||||
storage_key = f"generated/workspaces/{task.workspace_id}/projects/{task.project_id}/tasks/{task.id}/{output_name}"
|
||||
storage_key = (
|
||||
f"generated/workspaces/{task.workspace_id}/projects/{task.project_id}/tasks/{task.id}/{output_name}"
|
||||
)
|
||||
|
||||
with tempfile.TemporaryDirectory(prefix="xiaoxia-generation-") as temp_dir:
|
||||
temp_path = Path(temp_dir)
|
||||
@@ -235,7 +243,12 @@ def generate_video(task_id: str) -> dict:
|
||||
task.completed_at = datetime.now(timezone.utc)
|
||||
task_repo.update(task)
|
||||
|
||||
return {"status": "completed", "task_id": task.id, "video_id": video.id, "file_url": file_url}
|
||||
return {
|
||||
"status": "completed",
|
||||
"task_id": task.id,
|
||||
"video_id": video.id,
|
||||
"file_url": file_url,
|
||||
}
|
||||
except Exception as error:
|
||||
task.status = GenerationTaskStatus.FAILED
|
||||
task.error_message = str(error)
|
||||
|
||||
@@ -1,16 +1,20 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from packages.adapters.sqlalchemy_impl import SQLAlchemyAssetRepository, SQLAlchemyIngestJobRepository
|
||||
from packages.domain import Asset, IngestJobStatus
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl import (
|
||||
SQLAlchemyAssetRepository,
|
||||
SQLAlchemyIngestJobRepository,
|
||||
)
|
||||
from packages.domain import Asset, IngestJobStatus
|
||||
|
||||
|
||||
@celery_app.task(name="worker.ingest_asset")
|
||||
def ingest_asset(job_id: str) -> dict:
|
||||
"""
|
||||
Ingest asset task.
|
||||
|
||||
|
||||
Steps:
|
||||
1. Fetch IngestJob from repository
|
||||
2. Extract metadata from storage_key (placeholder: mock metadata)
|
||||
@@ -25,13 +29,13 @@ def ingest_asset(job_id: str) -> dict:
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
return {"status": "failed", "error": "job not found"}
|
||||
|
||||
|
||||
try:
|
||||
# Update job status to PROCESSING
|
||||
job.status = IngestJobStatus.PROCESSING
|
||||
job.updated_at = datetime.now(timezone.utc)
|
||||
job_repo.update(job)
|
||||
|
||||
|
||||
# Mock metadata extraction (in real implementation: use ffprobe, Pillow, etc.)
|
||||
mime_type = "video/mp4" if job.storage_key.endswith(".mp4") else "image/jpeg"
|
||||
metadata = {
|
||||
@@ -40,10 +44,10 @@ def ingest_asset(job_id: str) -> dict:
|
||||
"height": 1080,
|
||||
"size_bytes": 1024000,
|
||||
}
|
||||
|
||||
|
||||
# Extract filename from storage_key
|
||||
filename = job.storage_key.split("/")[-1]
|
||||
|
||||
|
||||
# Create Asset
|
||||
asset = Asset.create(
|
||||
workspace_id=job.workspace_id,
|
||||
@@ -55,7 +59,7 @@ def ingest_asset(job_id: str) -> dict:
|
||||
metadata=metadata,
|
||||
)
|
||||
asset_repo.create(asset)
|
||||
|
||||
|
||||
# Update job status to COMPLETED
|
||||
job.status = IngestJobStatus.COMPLETED
|
||||
job.result_asset_id = asset.id
|
||||
|
||||
Reference in New Issue
Block a user