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",
|
||||
|
||||
Reference in New Issue
Block a user