style: normalize python formatting gates

This commit is contained in:
Xiaoxia AI
2026-06-21 06:52:19 +08:00
parent 0809a079c5
commit bfbaddbd9a
129 changed files with 3024 additions and 2485 deletions
+1 -2
View File
@@ -1,8 +1,7 @@
import os
from logging.config import fileConfig
from sqlalchemy import engine_from_config
from sqlalchemy import pool
from sqlalchemy import engine_from_config, pool
from alembic import context
+85 -15
View File
@@ -12,9 +12,9 @@ run this migration normally.
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
from alembic import op
revision: str = "001"
down_revision: Union[str, None] = None
@@ -67,8 +67,18 @@ def upgrade() -> None:
sa.PrimaryKeyConstraint("id"),
)
op.create_index(op.f("ix_asset_libraries_kind"), "asset_libraries", ["kind"], unique=False)
op.create_index(op.f("ix_asset_libraries_project_id"), "asset_libraries", ["project_id"], unique=False)
op.create_index(op.f("ix_asset_libraries_workspace_id"), "asset_libraries", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_asset_libraries_project_id"),
"asset_libraries",
["project_id"],
unique=False,
)
op.create_index(
op.f("ix_asset_libraries_workspace_id"),
"asset_libraries",
["workspace_id"],
unique=False,
)
op.create_table(
"assets",
@@ -96,7 +106,12 @@ def upgrade() -> None:
sa.PrimaryKeyConstraint("id"),
)
op.create_index(op.f("ix_assets_asset_library_id"), "assets", ["asset_library_id"], unique=False)
op.create_index(op.f("ix_assets_classification_status"), "assets", ["classification_status"], unique=False)
op.create_index(
op.f("ix_assets_classification_status"),
"assets",
["classification_status"],
unique=False,
)
op.create_index(op.f("ix_assets_created_at"), "assets", ["created_at"], unique=False)
op.create_index(op.f("ix_assets_file_type"), "assets", ["file_type"], unique=False)
op.create_index(op.f("ix_assets_project_id"), "assets", ["project_id"], unique=False)
@@ -119,7 +134,12 @@ def upgrade() -> None:
)
op.create_index(op.f("ix_ingest_jobs_library_id"), "ingest_jobs", ["library_id"], unique=False)
op.create_index(op.f("ix_ingest_jobs_project_id"), "ingest_jobs", ["project_id"], unique=False)
op.create_index(op.f("ix_ingest_jobs_workspace_id"), "ingest_jobs", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_ingest_jobs_workspace_id"),
"ingest_jobs",
["workspace_id"],
unique=False,
)
op.create_table(
"classification_jobs",
@@ -135,9 +155,24 @@ def upgrade() -> None:
sa.Column("updated_at", sa.DateTime(), nullable=False),
sa.PrimaryKeyConstraint("id"),
)
op.create_index(op.f("ix_classification_jobs_asset_id"), "classification_jobs", ["asset_id"], unique=False)
op.create_index(op.f("ix_classification_jobs_project_id"), "classification_jobs", ["project_id"], unique=False)
op.create_index(op.f("ix_classification_jobs_workspace_id"), "classification_jobs", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_classification_jobs_asset_id"),
"classification_jobs",
["asset_id"],
unique=False,
)
op.create_index(
op.f("ix_classification_jobs_project_id"),
"classification_jobs",
["project_id"],
unique=False,
)
op.create_index(
op.f("ix_classification_jobs_workspace_id"),
"classification_jobs",
["workspace_id"],
unique=False,
)
op.create_table(
"generation_tasks",
@@ -157,10 +192,25 @@ def upgrade() -> None:
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.PrimaryKeyConstraint("id"),
)
op.create_index(op.f("ix_generation_tasks_asset_library_id"), "generation_tasks", ["asset_library_id"], unique=False)
op.create_index(op.f("ix_generation_tasks_project_id"), "generation_tasks", ["project_id"], unique=False)
op.create_index(
op.f("ix_generation_tasks_asset_library_id"),
"generation_tasks",
["asset_library_id"],
unique=False,
)
op.create_index(
op.f("ix_generation_tasks_project_id"),
"generation_tasks",
["project_id"],
unique=False,
)
op.create_index(op.f("ix_generation_tasks_status"), "generation_tasks", ["status"], unique=False)
op.create_index(op.f("ix_generation_tasks_workspace_id"), "generation_tasks", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_generation_tasks_workspace_id"),
"generation_tasks",
["workspace_id"],
unique=False,
)
op.create_table(
"generated_videos",
@@ -180,9 +230,24 @@ def upgrade() -> None:
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.PrimaryKeyConstraint("id"),
)
op.create_index(op.f("ix_generated_videos_generation_task_id"), "generated_videos", ["generation_task_id"], unique=False)
op.create_index(op.f("ix_generated_videos_project_id"), "generated_videos", ["project_id"], unique=False)
op.create_index(op.f("ix_generated_videos_workspace_id"), "generated_videos", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_generated_videos_generation_task_id"),
"generated_videos",
["generation_task_id"],
unique=False,
)
op.create_index(
op.f("ix_generated_videos_project_id"),
"generated_videos",
["project_id"],
unique=False,
)
op.create_index(
op.f("ix_generated_videos_workspace_id"),
"generated_videos",
["workspace_id"],
unique=False,
)
op.create_table(
"tasks",
@@ -244,7 +309,12 @@ def upgrade() -> None:
)
op.create_index(op.f("ix_task_issues_project_id"), "task_issues", ["project_id"], unique=False)
op.create_index(op.f("ix_task_issues_task_id"), "task_issues", ["task_id"], unique=False)
op.create_index(op.f("ix_task_issues_workspace_id"), "task_issues", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_task_issues_workspace_id"),
"task_issues",
["workspace_id"],
unique=False,
)
def downgrade() -> None:
+1 -2
View File
@@ -1,5 +1,3 @@
from fastapi import APIRouter
from app.api.routes.asset_libraries import router as asset_libraries_router
from app.api.routes.assets import router as assets_router
from app.api.routes.auth_simple import router as auth_router
@@ -11,6 +9,7 @@ from app.api.routes.ingest_jobs import router as ingest_jobs_router
from app.api.routes.project_management import router as project_management_router
from app.api.routes.projects import router as projects_router
from app.api.routes.upload import router as upload_router
from fastapi import APIRouter
api_router = APIRouter(prefix="/api/v1")
health_router = APIRouter()
+13 -4
View File
@@ -1,9 +1,18 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.dependencies import get_asset_library_repository
from app.schemas.asset_library import AssetLibraryResponse, CreateAssetLibraryRequest, ListAssetLibrariesResponse
from typing import Any
from packages.application import CreateAssetLibraryCommand, CreateAssetLibraryUseCase, ListAssetLibrariesUseCase
from app.schemas.asset_library import (
AssetLibraryResponse,
CreateAssetLibraryRequest,
ListAssetLibrariesResponse,
)
from fastapi import APIRouter, Depends
from packages.application import (
CreateAssetLibraryCommand,
CreateAssetLibraryUseCase,
ListAssetLibrariesUseCase,
)
from packages.domain import AssetLibraryKind
router = APIRouter()
+8 -3
View File
@@ -1,9 +1,14 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.dependencies import get_asset_repository
from app.schemas.asset import AssetResponse, CreateAssetRequest, ListAssetsResponse
from typing import Any
from packages.application import CreateAssetCommand, CreateAssetUseCase, ListAssetsUseCase
from fastapi import APIRouter, Depends
from packages.application import (
CreateAssetCommand,
CreateAssetUseCase,
ListAssetsUseCase,
)
from packages.domain import AssetStatus, ClassificationStatus
router = APIRouter()
+54 -50
View File
@@ -5,36 +5,35 @@ This module is intentionally disabled until the DI container/auth ports are rebu
Do not mount it directly; use `auth_simple.py` only as the current compatibility route.
"""
raise RuntimeError(
"apps.api.app.api.routes.auth is disabled: rebuild DI container before mounting full auth routes"
)
raise RuntimeError("apps.api.app.api.routes.auth is disabled: rebuild DI container before mounting full auth routes")
from fastapi import APIRouter, HTTPException, status, Depends, Request
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel, EmailStr
from packages.application.auth import (
RegisterUserUseCase,
RegisterUserRequest,
LoginUseCase,
LoginRequest,
LogoutUseCase,
LogoutRequest,
VerifyEmailUseCase,
VerifyEmailRequest,
RequestPasswordResetUseCase,
RequestPasswordResetRequest,
ResetPasswordUseCase,
ResetPasswordRequest,
)
from packages.domain.entities import User
from apps.api.app.dependencies import get_container
from apps.api.app.middleware.auth import get_current_user
from packages.application.auth import (
LoginRequest,
LoginUseCase,
LogoutRequest,
LogoutUseCase,
RegisterUserRequest,
RegisterUserUseCase,
RequestPasswordResetRequest,
RequestPasswordResetUseCase,
ResetPasswordRequest,
ResetPasswordUseCase,
VerifyEmailRequest,
VerifyEmailUseCase,
)
from packages.domain.entities import User
router = APIRouter(prefix="/auth", tags=["Authentication"])
# ==================== Request/Response Models ====================
class RegisterRequestModel(BaseModel):
email: EmailStr
password: str
@@ -77,11 +76,16 @@ class ResetPasswordModel(BaseModel):
# ==================== API Endpoints ====================
@router.post("/register", response_model=RegisterResponseModel, status_code=status.HTTP_201_CREATED)
@router.post(
"/register",
response_model=RegisterResponseModel,
status_code=status.HTTP_201_CREATED,
)
async def register(request: RegisterRequestModel):
"""
用户注册
- 邮箱必须唯一
- 用户名必须唯一
- 密码至少 8 位,包含大小写字母和数字
@@ -89,22 +93,22 @@ async def register(request: RegisterRequestModel):
"""
container = get_container()
use_case = container.get_register_user_use_case()
req = RegisterUserRequest(
email=request.email,
password=request.password,
username=request.username,
display_name=request.display_name,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=error,
)
return RegisterResponseModel(
user_id=response.user_id,
email=response.email,
@@ -118,7 +122,7 @@ async def register(request: RegisterRequestModel):
async def login(request: LoginRequestModel):
"""
用户登录
- 使用邮箱和密码登录
- 返回 access_token 和 refresh_token
- access_token 有效期 30 分钟
@@ -126,20 +130,20 @@ async def login(request: LoginRequestModel):
"""
container = get_container()
use_case = container.get_login_use_case()
req = LoginRequest(
email=request.email,
password=request.password,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=error,
)
return LoginResponseModel(
access_token=response.access_token,
refresh_token=response.refresh_token,
@@ -159,16 +163,16 @@ async def logout(
):
"""
用户登出
- 默认只登出当前设备
- 设置 logout_all_devices=true 可登出所有设备
"""
container = get_container()
use_case = container.get_logout_use_case()
# 从 JWT token 中提取 session_id
from packages.domain.auth import jwt_service
# 从 request 中获取 token
auth_header = request.headers.get("Authorization")
session_id = None
@@ -179,15 +183,15 @@ async def logout(
session_id = payload.get("sid") # 从 payload 提取 session_id
except:
pass # token 无效或没有 session_id,继续使用 None
req = LogoutRequest(
user_id=current_user.id,
session_id=session_id,
logout_all_devices=logout_all_devices,
)
success, error = use_case.execute(req)
if not success:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@@ -199,23 +203,23 @@ async def logout(
async def verify_email(token: str):
"""
邮箱验证
- 通过邮件中的链接访问此接口
- 验证成功后标记邮箱为已验证
"""
container = get_container()
use_case = container.get_verify_email_use_case()
req = VerifyEmailRequest(token=token)
success, error = use_case.execute(req)
if not success:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=error,
)
return {"message": "Email verified successfully"}
@@ -223,18 +227,18 @@ async def verify_email(token: str):
async def forgot_password(request: PasswordResetRequestModel):
"""
请求密码重置
- 发送密码重置邮件
- 邮件中包含重置链接(有效期 1 小时)
- 即使邮箱不存在也返回成功(安全考虑)
"""
container = get_container()
use_case = container.get_request_password_reset_use_case()
req = RequestPasswordResetRequest(email=request.email)
success, error = use_case.execute(req)
# 不论成功失败都返回 202(安全考虑)
return {"message": "Password reset email sent if account exists"}
@@ -243,24 +247,24 @@ async def forgot_password(request: PasswordResetRequestModel):
async def reset_password(request: ResetPasswordModel):
"""
重置密码
- 使用邮件中的 token 重置密码
- 新密码必须符合密码强度要求
"""
container = get_container()
use_case = container.get_reset_password_use_case()
req = ResetPasswordRequest(
token=request.token,
new_password=request.new_password,
)
success, error = use_case.execute(req)
if not success:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=error,
)
return {"message": "Password reset successfully"}
+8 -4
View File
@@ -1,17 +1,18 @@
"""
认证 APISQLAlchemy ORM
"""
from datetime import datetime, timedelta, timezone
import hashlib
import secrets
from datetime import datetime, timedelta, timezone
import jwt
from app.config import settings
from app.dependencies import get_db_session
from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel, EmailStr
from sqlalchemy.orm import Session
from app.config import settings
from app.dependencies import get_db_session
from packages.adapters.sqlalchemy_impl.models import UserModel
from packages.domain.auth import password_hasher, password_validator
@@ -169,4 +170,7 @@ async def login(request: LoginRequest, db: Session = Depends(get_db_session)):
@router.get("/me")
async def get_current_user_info():
raise HTTPException(status_code=status.HTTP_501_NOT_IMPLEMENTED, detail="/auth/me requires bearer-token dependency integration")
raise HTTPException(
status_code=status.HTTP_501_NOT_IMPLEMENTED,
detail="/auth/me requires bearer-token dependency integration",
)
+11 -5
View File
@@ -1,12 +1,18 @@
from datetime import datetime, timezone
from fastapi import APIRouter, Depends, HTTPException
from typing import Any
from app.core.celery_app import celery_app
from app.dependencies import get_classification_job_repository
from app.schemas.classification_job import ClassificationJobResponse, SubmitClassificationJobRequest
from typing import Any
from packages.application import SubmitClassificationJobCommand, SubmitClassificationJobUseCase
from app.schemas.classification_job import (
ClassificationJobResponse,
SubmitClassificationJobRequest,
)
from fastapi import APIRouter, Depends, HTTPException
from packages.application import (
SubmitClassificationJobCommand,
SubmitClassificationJobUseCase,
)
router = APIRouter()
+3 -2
View File
@@ -1,4 +1,4 @@
from fastapi import APIRouter, Depends, HTTPException
from typing import Any
from app.core.storage import MinIOService, get_minio_service
from app.dependencies import get_generated_video_repository
@@ -7,7 +7,8 @@ from app.schemas.generated_video import (
GeneratedVideoResponse,
ListGeneratedVideosResponse,
)
from typing import Any
from fastapi import APIRouter, Depends, HTTPException
from packages.application import (
GetGeneratedVideoDownloadUrlUseCase,
GetGeneratedVideoUseCase,
+15 -5
View File
@@ -1,10 +1,20 @@
from fastapi import APIRouter, Depends, HTTPException
from typing import Any
from app.core.celery_app import celery_app
from app.dependencies import get_generation_task_repository, get_generated_video_repository
from app.schemas.generation_task import CreateGenerationTaskRequest, GenerationTaskResponse
from app.schemas.generated_video import GeneratedVideoResponse, ListGeneratedVideosResponse
from typing import Any
from app.dependencies import (
get_generated_video_repository,
get_generation_task_repository,
)
from app.schemas.generated_video import (
GeneratedVideoResponse,
ListGeneratedVideosResponse,
)
from app.schemas.generation_task import (
CreateGenerationTaskRequest,
GenerationTaskResponse,
)
from fastapi import APIRouter, Depends, HTTPException
from packages.application import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
+33 -13
View File
@@ -1,12 +1,11 @@
from pydantic import BaseModel
from datetime import datetime
import psycopg2
import redis
from app.config import settings
from fastapi import APIRouter, status
from fastapi.responses import JSONResponse
from app.config import settings
from pydantic import BaseModel
router = APIRouter(tags=["Health"])
@@ -56,16 +55,28 @@ async def startup_check():
async def _check_database() -> dict:
if settings.USE_IN_MEMORY_DB:
return {"status": "healthy", "type": "in_memory", "message": "Using in-memory database"}
return {
"status": "healthy",
"type": "in_memory",
"message": "Using in-memory database",
}
try:
conn = psycopg2.connect(settings.DATABASE_URL, connect_timeout=3)
with conn.cursor() as cur:
cur.execute("SELECT 1")
cur.fetchone()
conn.close()
return {"status": "healthy", "type": "postgresql", "message": "Database connection successful"}
return {
"status": "healthy",
"type": "postgresql",
"message": "Database connection successful",
}
except Exception as error:
return {"status": "unhealthy", "type": "postgresql", "message": f"Database connection failed: {error}"}
return {
"status": "unhealthy",
"type": "postgresql",
"message": f"Database connection failed: {error}",
}
async def _check_redis() -> dict:
@@ -73,23 +84,32 @@ async def _check_redis() -> dict:
client = redis.from_url(settings.REDIS_URL, socket_connect_timeout=3)
client.ping()
client.close()
return {"status": "healthy", "type": "redis", "message": "Redis connection successful"}
return {
"status": "healthy",
"type": "redis",
"message": "Redis connection successful",
}
except Exception as error:
return {"status": "unhealthy", "type": "redis", "message": f"Redis connection failed: {error}"}
return {
"status": "unhealthy",
"type": "redis",
"message": f"Redis connection failed: {error}",
}
async def _check_migrations() -> dict:
if settings.USE_IN_MEMORY_DB:
return {"status": "healthy", "message": "Using in-memory database, no migrations needed"}
return {
"status": "healthy",
"message": "Using in-memory database, no migrations needed",
}
try:
conn = psycopg2.connect(settings.DATABASE_URL, connect_timeout=3)
with conn.cursor() as cur:
cur.execute(
"""
cur.execute("""
SELECT COUNT(*) FROM information_schema.tables
WHERE table_name IN ('projects', 'asset_libraries', 'assets', 'ingest_jobs', 'classification_jobs')
"""
)
""")
count = cur.fetchone()[0]
conn.close()
if count >= 5:
+3 -2
View File
@@ -1,9 +1,10 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.core.celery_app import celery_app
from app.dependencies import get_ingest_job_repository
from app.schemas.ingest_job import IngestJobResponse, SubmitIngestJobRequest
from typing import Any
from fastapi import APIRouter, Depends
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
router = APIRouter()
@@ -1,4 +1,5 @@
"""项目管理 API 路由"""
from datetime import datetime
from typing import Annotated
@@ -11,7 +12,6 @@ from packages.adapters.sqlite_tracker.project_management_repositories import (
SQLiteTaskRepository,
)
from packages.application.get_task_detail_use_case import GetTaskDetailUseCase
from packages.application.update_task_use_case import UpdateTaskUseCase
from packages.application.project_management_use_cases import (
CreateMilestoneUseCase,
CreateTaskIssueUseCase,
@@ -23,6 +23,7 @@ from packages.application.project_management_use_cases import (
UpdateTaskProgressUseCase,
UpdateTaskStatusUseCase,
)
from packages.application.update_task_use_case import UpdateTaskUseCase
from packages.domain import TaskPriority, TaskStatus
router = APIRouter()
+13 -4
View File
@@ -1,9 +1,18 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.dependencies import get_project_repository
from app.schemas.project import CreateProjectRequest, ListProjectsResponse, ProjectResponse
from typing import Any
from packages.application import CreateProjectCommand, CreateProjectUseCase, ListProjectsUseCase
from app.schemas.project import (
CreateProjectRequest,
ListProjectsResponse,
ProjectResponse,
)
from fastapi import APIRouter, Depends
from packages.application import (
CreateProjectCommand,
CreateProjectUseCase,
ListProjectsUseCase,
)
router = APIRouter()
+3 -2
View File
@@ -1,11 +1,12 @@
from fastapi import APIRouter, Depends, File, Form, UploadFile
from typing import Any
from uuid import uuid4
from app.core.celery_app import celery_app
from app.core.storage import MinIOService, get_minio_service
from app.dependencies import get_ingest_job_repository
from app.schemas.upload import UploadAssetResponse
from typing import Any
from fastapi import APIRouter, Depends, File, Form, UploadFile
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
router = APIRouter()
+60 -49
View File
@@ -9,21 +9,28 @@ raise RuntimeError(
"apps.api.app.api.routes.workspaces is disabled: rebuild DI container before mounting workspace routes"
)
from fastapi import APIRouter, HTTPException, status, Depends
from pydantic import BaseModel, EmailStr
from typing import List
from datetime import datetime
from typing import List
from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel, EmailStr
from apps.api.app.dependencies import get_container
from apps.api.app.middleware.auth import (
get_current_user,
require_workspace_access,
require_workspace_admin,
require_workspace_owner,
)
from packages.application.workspace import *
from packages.domain.entities import User
from apps.api.app.dependencies import get_container
from apps.api.app.middleware.auth import get_current_user, require_workspace_access, require_workspace_admin, require_workspace_owner
router = APIRouter(prefix="/workspaces", tags=["Workspaces"])
# ==================== Request/Response Models ====================
class CreateWorkspaceRequestModel(BaseModel):
name: str
subscription_plan: str = "free"
@@ -52,6 +59,7 @@ class UpgradeSubscriptionRequestModel(BaseModel):
# ==================== Workspace CRUD ====================
@router.post("", response_model=WorkspaceResponseModel, status_code=status.HTTP_201_CREATED)
async def create_workspace(
request: CreateWorkspaceRequestModel,
@@ -60,18 +68,18 @@ async def create_workspace(
"""创建工作空间"""
container = get_container()
use_case = container.get_create_workspace_use_case()
req = CreateWorkspaceRequest(
name=request.name,
owner_user_id=current_user.id,
subscription_plan=request.subscription_plan,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return WorkspaceResponseModel(
workspace_id=response.workspace_id,
name=response.name,
@@ -86,13 +94,13 @@ async def list_workspaces(current_user: User = Depends(get_current_user)):
"""获取用户的所有工作空间"""
container = get_container()
use_case = container.get_list_workspaces_use_case()
req = ListWorkspacesRequest(user_id=current_user.id)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {
"workspaces": [
{
@@ -117,13 +125,13 @@ async def get_workspace_detail(
"""获取工作空间详情"""
container = get_container()
use_case = container.get_get_workspace_detail_use_case()
req = GetWorkspaceDetailRequest(workspace_id=workspace_id, user_id=current_user.id)
detail, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=error)
return {
"workspace_id": detail.workspace_id,
"name": detail.name,
@@ -140,6 +148,7 @@ async def get_workspace_detail(
# ==================== Member Management ====================
@router.post("/{workspace_id}/members/invite", status_code=status.HTTP_201_CREATED)
async def invite_member(
workspace_id: str,
@@ -149,19 +158,19 @@ async def invite_member(
"""邀请成员"""
container = get_container()
use_case = container.get_invite_member_use_case()
req = InviteMemberRequest(
workspace_id=workspace_id,
inviter_user_id=current_user.id,
invitee_email=request.email,
role=request.role,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {
"invitation_id": response.invitation_id,
"invitee_email": response.invitee_email,
@@ -178,13 +187,13 @@ async def list_members(
"""获取成员列表"""
container = get_container()
use_case = container.get_list_members_use_case()
req = ListMembersRequest(workspace_id=workspace_id, requester_user_id=current_user.id)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=error)
return {
"members": [
{
@@ -211,15 +220,15 @@ async def remove_member(
"""移除成员"""
container = get_container()
use_case = container.get_remove_member_use_case()
req = RemoveMemberRequest(
workspace_id=workspace_id,
requester_user_id=current_user.id,
target_user_id=user_id,
)
success, error = use_case.execute(req)
if not success:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
@@ -232,10 +241,10 @@ async def leave_workspace(
"""离开工作空间"""
container = get_container()
use_case = container.get_leave_workspace_use_case()
req = LeaveWorkspaceRequest(workspace_id=workspace_id, user_id=current_user.id)
success, error = use_case.execute(req)
if not success:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
@@ -250,19 +259,19 @@ async def update_member_role(
"""修改成员角色"""
container = get_container()
use_case = container.get_update_member_role_use_case()
req = UpdateMemberRoleRequest(
workspace_id=workspace_id,
requester_user_id=current_user.id,
target_user_id=user_id,
new_role=request.role,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {
"user_id": response.user_id,
"old_role": response.old_role,
@@ -272,6 +281,7 @@ async def update_member_role(
# ==================== Subscription Management ====================
@router.post("/{workspace_id}/subscription/upgrade")
async def upgrade_subscription(
workspace_id: str,
@@ -281,18 +291,18 @@ async def upgrade_subscription(
"""升级订阅"""
container = get_container()
use_case = container.get_upgrade_subscription_use_case()
req = UpgradeSubscriptionRequest(
workspace_id=workspace_id,
requester_user_id=current_user.id,
new_plan=request.new_plan,
)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {
"workspace_id": response.workspace_id,
"old_plan": response.old_plan,
@@ -310,13 +320,13 @@ async def cancel_subscription(
"""取消订阅"""
container = get_container()
use_case = container.get_cancel_subscription_use_case()
req = CancelSubscriptionRequest(workspace_id=workspace_id, requester_user_id=current_user.id)
success, error = use_case.execute(req)
if not success:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {"message": "Subscription cancelled successfully"}
@@ -328,24 +338,25 @@ async def get_quota_status(
"""获取配额状态"""
container = get_container()
quota_checker = container.quota_checker
# 检查权限
permission_checker = container.permission_checker
has_access, _ = permission_checker.check_workspace_access(workspace_id, current_user.id)
if not has_access:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Access denied")
status = quota_checker.get_quota_status(workspace_id)
if not status:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Workspace not found")
return status
# ==================== Invitation Acceptance ====================
@router.post("/invitations/{token}/accept")
async def accept_invitation(
token: str,
@@ -354,13 +365,13 @@ async def accept_invitation(
"""接受邀请"""
container = get_container()
use_case = container.get_accept_invitation_use_case()
req = AcceptInvitationRequest(invitation_token=token, user_id=current_user.id)
response, error = use_case.execute(req)
if error:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {
"workspace_id": response.workspace_id,
"workspace_name": response.workspace_name,
@@ -373,11 +384,11 @@ async def decline_invitation(token: str):
"""拒绝邀请"""
container = get_container()
use_case = container.get_decline_invitation_use_case()
req = DeclineInvitationRequest(invitation_token=token)
success, error = use_case.execute(req)
if not success:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error)
return {"message": "Invitation declined"}
+3 -2
View File
@@ -1,6 +1,7 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
from typing import Optional
import os
from typing import Optional
from pydantic_settings import BaseSettings, SettingsConfigDict
class Settings(BaseSettings):
+1 -3
View File
@@ -1,7 +1,5 @@
from celery import Celery
from app.config import get_settings
from celery import Celery
settings = get_settings()
celery_app = Celery("xiaoxia-saas-api")
+3 -2
View File
@@ -1,6 +1,7 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
from typing import Optional
import os
from typing import Optional
from pydantic_settings import BaseSettings, SettingsConfigDict
class DatabaseSettings(BaseSettings):
+18 -22
View File
@@ -1,7 +1,8 @@
"""阿里云 OSS 存储服务"""
import logging
from urllib.parse import urlparse
import os
from urllib.parse import urlparse
try:
import oss2
@@ -48,12 +49,12 @@ class OSSStorageService:
) -> str:
"""
上传文件到 OSS
Args:
file_or_path: 文件对象或本地文件路径
storage_key: 存储键(文件路径)
content_type: 内容类型
Returns:
文件公网 URL
"""
@@ -63,20 +64,12 @@ class OSSStorageService:
try:
# 如果是字符串路径,从本地文件上传
if isinstance(file_or_path, str):
self.bucket.put_object_from_file(
storage_key,
file_or_path,
headers={'Content-Type': content_type}
)
self.bucket.put_object_from_file(storage_key, file_or_path, headers={"Content-Type": content_type})
else:
# 文件对象
file_or_path.seek(0)
self.bucket.put_object(
storage_key,
file_or_path,
headers={'Content-Type': content_type}
)
self.bucket.put_object(storage_key, file_or_path, headers={"Content-Type": content_type})
return f"{self.public_url}/{storage_key}"
except Exception as e:
raise Exception(f"Failed to upload file to OSS: {e}")
@@ -88,11 +81,11 @@ class OSSStorageService:
def get_download_url(self, storage_key_or_url: str, expires_seconds: int = 3600) -> str:
"""
获取文件下载签名 URL(用于私有文件)
Args:
storage_key_or_url: 存储键或完整 URL
expires_seconds: 过期时间(秒)
Returns:
签名 URL
"""
@@ -103,7 +96,7 @@ class OSSStorageService:
storage_key = self._normalize_storage_key(storage_key_or_url)
try:
return self.bucket.sign_url('GET', storage_key, expires_seconds)
return self.bucket.sign_url("GET", storage_key, expires_seconds)
except Exception:
return self.get_url(storage_key)
@@ -118,7 +111,7 @@ class OSSStorageService:
def download_file(self, storage_key: str, local_path: str):
"""
从 OSS 下载文件到本地
Args:
storage_key: 存储键
local_path: 本地文件路径
@@ -135,7 +128,7 @@ class OSSStorageService:
def delete_file(self, storage_key: str):
"""
删除 OSS 文件
Args:
storage_key: 存储键
"""
@@ -145,15 +138,18 @@ class OSSStorageService:
try:
self.bucket.delete_object(storage_key)
except Exception as error:
logger.warning("Failed to delete file from OSS", extra={"storage_key": storage_key, "error": str(error)})
logger.warning(
"Failed to delete file from OSS",
extra={"storage_key": storage_key, "error": str(error)},
)
def file_exists(self, storage_key: str) -> bool:
"""
检查文件是否存在
Args:
storage_key: 存储键
Returns:
是否存在
"""
+6 -2
View File
@@ -1,9 +1,13 @@
from collections.abc import Generator
from app.config import settings
from sqlalchemy.orm import Session
from app.config import settings
from packages.adapters.sqlalchemy_impl import build_session_factory, ensure_database_exists, initialize_database
from packages.adapters.sqlalchemy_impl import (
build_session_factory,
ensure_database_exists,
initialize_database,
)
ensure_database_exists(settings.DATABASE_URL)
engine, SessionLocal = build_session_factory(
+40 -14
View File
@@ -1,14 +1,26 @@
from app.config import settings
from fastapi import Depends
from sqlalchemy.orm import Session
from app.config import settings
from packages.adapters.sqlalchemy_impl.asset_library_repository import SQLAlchemyAssetLibraryRepository
from packages.adapters.sqlalchemy_impl.asset_library_repository import (
SQLAlchemyAssetLibraryRepository,
)
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from packages.adapters.sqlalchemy_impl.classification_job_repository import SQLAlchemyClassificationJobRepository
from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository
from packages.adapters.sqlalchemy_impl.generation_task_repository import SQLAlchemyGenerationTaskRepository
from packages.adapters.sqlalchemy_impl.ingest_job_repository import SQLAlchemyIngestJobRepository
from packages.adapters.sqlalchemy_impl.project_repository import SQLAlchemyProjectRepository
from packages.adapters.sqlalchemy_impl.classification_job_repository import (
SQLAlchemyClassificationJobRepository,
)
from packages.adapters.sqlalchemy_impl.generated_video_repository import (
SQLAlchemyGeneratedVideoRepository,
)
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
from packages.adapters.sqlalchemy_impl.ingest_job_repository import (
SQLAlchemyIngestJobRepository,
)
from packages.adapters.sqlalchemy_impl.project_repository import (
SQLAlchemyProjectRepository,
)
from packages.adapters.sqlalchemy_impl.session import build_session_factory
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
@@ -22,29 +34,43 @@ def get_db_session():
session.close()
def get_asset_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyAssetRepository:
def get_asset_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyAssetRepository:
return SQLAlchemyAssetRepository(session)
def get_asset_library_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyAssetLibraryRepository:
def get_asset_library_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyAssetLibraryRepository:
return SQLAlchemyAssetLibraryRepository(session)
def get_ingest_job_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyIngestJobRepository:
def get_ingest_job_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyIngestJobRepository:
return SQLAlchemyIngestJobRepository(session)
def get_classification_job_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyClassificationJobRepository:
def get_classification_job_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyClassificationJobRepository:
return SQLAlchemyClassificationJobRepository(session)
def get_generation_task_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyGenerationTaskRepository:
def get_generation_task_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyGenerationTaskRepository:
return SQLAlchemyGenerationTaskRepository(session)
def get_generated_video_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyGeneratedVideoRepository:
def get_generated_video_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyGeneratedVideoRepository:
return SQLAlchemyGeneratedVideoRepository(session)
def get_project_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyProjectRepository:
def get_project_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyProjectRepository:
return SQLAlchemyProjectRepository(session)
+31 -34
View File
@@ -5,17 +5,14 @@ Disabled because it depends on the removed DI container. Rebuild it around the
canonical JWT settings and SQLAlchemy-backed user repository before reuse.
"""
raise RuntimeError(
"apps.api.app.middleware.auth is disabled: rebuild auth dependency wiring before importing it"
)
raise RuntimeError("apps.api.app.middleware.auth is disabled: rebuild auth dependency wiring before importing it")
from fastapi import Depends, HTTPException, status
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from apps.api.app.dependencies import get_container
from packages.domain.auth import jwt_service
from packages.domain.entities import User
from apps.api.app.dependencies import get_container
security = HTTPBearer()
@@ -25,42 +22,42 @@ async def get_current_user(
) -> User:
"""
获取当前登录用户
从 Authorization header 中提取 JWT token 并验证
Raises:
HTTPException: Token 无效或过期
Returns:
当前用户对象
"""
token = credentials.credentials
try:
# 验证 token
payload = jwt_service.verify_token(token)
user_id = payload.get("sub")
if not user_id:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid token: missing user_id",
headers={"WWW-Authenticate": "Bearer"},
)
# 从数据库获取用户
container = get_container()
user = container.user_repository.find_by_id(user_id)
if not user:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="User not found",
headers={"WWW-Authenticate": "Bearer"},
)
return user
except Exception as e:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
@@ -74,15 +71,15 @@ async def get_current_user_optional(
) -> User | None:
"""
获取当前登录用户(可选)
如果没有提供 token,返回 None 而不是抛出异常
Returns:
当前用户对象或 None
"""
if not credentials:
return None
try:
return await get_current_user(credentials)
except HTTPException:
@@ -92,78 +89,78 @@ async def get_current_user_optional(
def require_workspace_access(workspace_id: str, user: User = Depends(get_current_user)) -> tuple[str, str]:
"""
要求用户可以访问指定工作空间
Args:
workspace_id: 工作空间 ID
user: 当前用户
Raises:
HTTPException: 用户没有访问权限
Returns:
(workspace_id, user_role)
"""
container = get_container()
permission_checker = container.permission_checker
has_access, role = permission_checker.check_workspace_access(workspace_id, user.id)
if not has_access:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="You don't have access to this workspace",
)
return workspace_id, role
def require_workspace_admin(workspace_id: str, user: User = Depends(get_current_user)) -> str:
"""
要求用户是工作空间的 Admin 或 Owner
Args:
workspace_id: 工作空间 ID
user: 当前用户
Raises:
HTTPException: 用户没有管理权限
Returns:
workspace_id
"""
container = get_container()
permission_checker = container.permission_checker
if not permission_checker.check_is_admin_or_owner(workspace_id, user.id):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only workspace owner or admin can perform this action",
)
return workspace_id
def require_workspace_owner(workspace_id: str, user: User = Depends(get_current_user)) -> str:
"""
要求用户是工作空间的 Owner
Args:
workspace_id: 工作空间 ID
user: 当前用户
Raises:
HTTPException: 用户不是 Owner
Returns:
workspace_id
"""
container = get_container()
permission_checker = container.permission_checker
if not permission_checker.check_is_owner(workspace_id, user.id):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Only workspace owner can perform this action",
)
return workspace_id
+21 -17
View File
@@ -1,19 +1,21 @@
"""
全局异常处理和错误响应
"""
from fastapi import Request, status
from fastapi.responses import JSONResponse
from fastapi.exceptions import RequestValidationError
from starlette.exceptions import HTTPException as StarletteHTTPException
import traceback
import logging
import traceback
from fastapi import Request, status
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse
from starlette.exceptions import HTTPException as StarletteHTTPException
logger = logging.getLogger(__name__)
class APIException(Exception):
"""API 异常基类"""
def __init__(
self,
message: str,
@@ -28,7 +30,7 @@ class APIException(Exception):
class AuthenticationError(APIException):
"""认证错误"""
def __init__(self, message: str = "Authentication failed"):
super().__init__(
message=message,
@@ -39,7 +41,7 @@ class AuthenticationError(APIException):
class PermissionDeniedError(APIException):
"""权限拒绝"""
def __init__(self, message: str = "Permission denied"):
super().__init__(
message=message,
@@ -50,7 +52,7 @@ class PermissionDeniedError(APIException):
class ResourceNotFoundError(APIException):
"""资源不存在"""
def __init__(self, resource: str = "Resource"):
super().__init__(
message=f"{resource} not found",
@@ -61,7 +63,7 @@ class ResourceNotFoundError(APIException):
class ValidationError(APIException):
"""验证错误"""
def __init__(self, message: str):
super().__init__(
message=message,
@@ -100,12 +102,14 @@ async def validation_exception_handler(request: Request, exc: RequestValidationE
"""请求验证异常处理"""
errors = []
for error in exc.errors():
errors.append({
"field": ".".join(str(loc) for loc in error["loc"]),
"message": error["msg"],
"type": error["type"],
})
errors.append(
{
"field": ".".join(str(loc) for loc in error["loc"]),
"message": error["msg"],
"type": error["type"],
}
)
return JSONResponse(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
content={
@@ -121,7 +125,7 @@ async def validation_exception_handler(request: Request, exc: RequestValidationE
async def general_exception_handler(request: Request, exc: Exception):
"""通用异常处理"""
logger.error(f"Unhandled exception: {exc}", exc_info=True)
# 生产环境不返回详细错误信息
return JSONResponse(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
+23 -24
View File
@@ -1,8 +1,10 @@
"""
请求日志中间件
"""
import time
import logging
import time
from fastapi import Request
from starlette.middleware.base import BaseHTTPMiddleware
@@ -11,58 +13,57 @@ logger = logging.getLogger(__name__)
class RequestLoggingMiddleware(BaseHTTPMiddleware):
"""请求日志中间件"""
async def dispatch(self, request: Request, call_next):
# 记录请求开始时间
start_time = time.time()
# 记录请求信息
logger.info(f"Request: {request.method} {request.url.path}")
# 处理请求
response = await call_next(request)
# 计算处理时间
process_time = time.time() - start_time
# 记录响应信息
logger.info(
f"Response: {request.method} {request.url.path} "
f"status={response.status_code} time={process_time:.3f}s"
f"Response: {request.method} {request.url.path} " f"status={response.status_code} time={process_time:.3f}s"
)
# 添加响应头
response.headers["X-Process-Time"] = str(process_time)
return response
class RateLimitMiddleware(BaseHTTPMiddleware):
"""简单的速率限制中间件(基于内存)"""
def __init__(self, app, max_requests: int = 100, window_seconds: int = 60):
super().__init__(app)
self.max_requests = max_requests
self.window_seconds = window_seconds
self.requests = {} # {ip: [(timestamp, ...)]}
async def dispatch(self, request: Request, call_next):
# 获取客户端 IP
client_ip = request.client.host
current_time = time.time()
# 清理过期记录
if client_ip in self.requests:
self.requests[client_ip] = [
ts for ts in self.requests[client_ip]
if current_time - ts < self.window_seconds
ts for ts in self.requests[client_ip] if current_time - ts < self.window_seconds
]
# 检查速率限制
request_count = len(self.requests.get(client_ip, []))
if request_count >= self.max_requests:
from fastapi.responses import JSONResponse
return JSONResponse(
status_code=429,
content={
@@ -72,19 +73,17 @@ class RateLimitMiddleware(BaseHTTPMiddleware):
}
},
)
# 记录请求
if client_ip not in self.requests:
self.requests[client_ip] = []
self.requests[client_ip].append(current_time)
# 处理请求
response = await call_next(request)
# 添加速率限制信息到响应头
response.headers["X-RateLimit-Limit"] = str(self.max_requests)
response.headers["X-RateLimit-Remaining"] = str(
self.max_requests - len(self.requests[client_ip])
)
response.headers["X-RateLimit-Remaining"] = str(self.max_requests - len(self.requests[client_ip]))
return response
+28 -26
View File
@@ -1,9 +1,11 @@
"""
性能监控中间件
"""
import time
import logging
import time
from typing import Callable
from fastapi import Request, Response
from starlette.middleware.base import BaseHTTPMiddleware
@@ -12,30 +14,30 @@ logger = logging.getLogger(__name__)
class PerformanceMonitoringMiddleware(BaseHTTPMiddleware):
"""性能监控中间件"""
def __init__(self, app, slow_request_threshold: float = 1.0):
super().__init__(app)
self.slow_request_threshold = slow_request_threshold # 慢请求阈值(秒)
async def dispatch(self, request: Request, call_next: Callable):
# 记录请求开始时间
start_time = time.time()
# 生成请求 ID
request_id = self._generate_request_id()
request.state.request_id = request_id
# 处理请求
try:
response = await call_next(request)
# 计算处理时间
process_time = time.time() - start_time
# 添加响应头
response.headers["X-Request-ID"] = request_id
response.headers["X-Process-Time"] = f"{process_time:.3f}"
# 记录慢请求
if process_time > self.slow_request_threshold:
logger.warning(
@@ -43,55 +45,55 @@ class PerformanceMonitoringMiddleware(BaseHTTPMiddleware):
f"took {process_time:.3f}s (threshold: {self.slow_request_threshold}s) "
f"[request_id={request_id}]"
)
# 记录请求日志
logger.info(
f"{request.method} {request.url.path} "
f"status={response.status_code} time={process_time:.3f}s "
f"[request_id={request_id}]"
)
return response
except Exception as e:
process_time = time.time() - start_time
logger.error(
f"Request failed: {request.method} {request.url.path} "
f"error={str(e)} time={process_time:.3f}s "
f"[request_id={request_id}]",
exc_info=True
exc_info=True,
)
raise
def _generate_request_id(self) -> str:
"""生成请求 ID"""
import uuid
return str(uuid.uuid4())
class DatabaseQueryLogger:
"""数据库查询日志记录器"""
def __init__(self):
self.queries = []
self.total_time = 0
def log_query(self, query: str, params: tuple, duration: float):
"""记录查询"""
self.queries.append({
"query": query,
"params": params,
"duration": duration,
})
self.queries.append(
{
"query": query,
"params": params,
"duration": duration,
}
)
self.total_time += duration
# 记录慢查询(超过 100ms
if duration > 0.1:
logger.warning(
f"Slow query detected: {query[:100]}... "
f"took {duration:.3f}s with params {params}"
)
logger.warning(f"Slow query detected: {query[:100]}... " f"took {duration:.3f}s with params {params}")
def get_stats(self):
"""获取统计信息"""
return {
+21 -20
View File
@@ -1,14 +1,16 @@
"""
API 版本管理中间件
"""
from datetime import datetime
from fastapi import Request
from starlette.middleware.base import BaseHTTPMiddleware
from datetime import datetime
class APIVersionMiddleware(BaseHTTPMiddleware):
"""API 版本管理中间件"""
# 版本配置
VERSIONS = {
"v1": {
@@ -24,33 +26,31 @@ class APIVersionMiddleware(BaseHTTPMiddleware):
"release_date": None,
},
}
async def dispatch(self, request: Request, call_next):
# 提取版本号
version = self._extract_version(request.url.path)
# 处理请求
response = await call_next(request)
# 添加版本信息头
if version:
response.headers["X-API-Version"] = version
# 添加弃用警告
version_info = self.VERSIONS.get(version, {})
if version_info.get("deprecated"):
response.headers["X-API-Deprecated"] = "true"
sunset_date = version_info.get("sunset_date")
if sunset_date:
response.headers["X-API-Sunset-Date"] = sunset_date
response.headers["X-API-Deprecation-Info"] = (
f"https://docs.xiaoxia-saas.com/api/deprecation/{version}"
)
response.headers["X-API-Deprecation-Info"] = f"https://docs.xiaoxia-saas.com/api/deprecation/{version}"
return response
def _extract_version(self, path: str) -> str:
"""从路径中提取版本号"""
parts = path.split("/")
@@ -62,14 +62,15 @@ class APIVersionMiddleware(BaseHTTPMiddleware):
class VersionNotFoundMiddleware(BaseHTTPMiddleware):
"""处理已下线的 API 版本"""
SUNSET_VERSIONS = [] # 已下线的版本列表
async def dispatch(self, request: Request, call_next):
version = self._extract_version(request.url.path)
if version in self.SUNSET_VERSIONS:
from fastapi.responses import JSONResponse
return JSONResponse(
status_code=410,
content={
@@ -77,13 +78,13 @@ class VersionNotFoundMiddleware(BaseHTTPMiddleware):
"code": "API_VERSION_SUNSET",
"message": f"API {version} has been sunset and is no longer available",
"sunset_date": "2028-07-01",
"migration_guide": f"https://docs.xiaoxia-saas.com/api/migration/{version}"
"migration_guide": f"https://docs.xiaoxia-saas.com/api/migration/{version}",
}
}
},
)
return await call_next(request)
def _extract_version(self, path: str) -> str:
"""从路径中提取版本号"""
parts = path.split("/")
+5 -1
View File
@@ -1,7 +1,11 @@
"""Schema package."""
from .asset import AssetResponse, CreateAssetRequest, ListAssetsResponse
from .asset_library import AssetLibraryResponse, CreateAssetLibraryRequest, ListAssetLibrariesResponse
from .asset_library import (
AssetLibraryResponse,
CreateAssetLibraryRequest,
ListAssetLibrariesResponse,
)
from .health import HealthResponse
from .ingest_job import IngestJobResponse, SubmitIngestJobRequest
from .project import CreateProjectRequest, ListProjectsResponse, ProjectResponse
+6 -7
View File
@@ -1,12 +1,5 @@
import os
from fastapi import FastAPI
from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from fastapi.middleware.gzip import GZipMiddleware
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.staticfiles import StaticFiles
from app.api.router import api_router, health_router
from app.config import settings
from app.middleware.exceptions import (
@@ -17,6 +10,12 @@ from app.middleware.exceptions import (
validation_exception_handler,
)
from app.middleware.logging import RequestLoggingMiddleware
from fastapi import FastAPI
from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from fastapi.middleware.gzip import GZipMiddleware
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.staticfiles import StaticFiles
app = FastAPI(
title="小虾 SaaS API",
+28 -18
View File
@@ -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
View File
@@ -1,6 +1,7 @@
"""
视频处理模块
"""
from .processor import VideoProcessor, VideoResult
__all__ = ["VideoProcessor", "VideoResult"]
+33 -33
View File
@@ -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
-2
View File
@@ -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
+3 -2
View File
@@ -1,6 +1,7 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
from typing import Optional
import os
from typing import Optional
from pydantic_settings import BaseSettings, SettingsConfigDict
class WorkerSettings(BaseSettings):
+6 -1
View File
@@ -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)
+16 -9
View File
@@ -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,
+19 -6
View File
@@ -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)
+12 -8
View File
@@ -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
@@ -1,4 +1,5 @@
"""AssetLibrary InMemory Repository 实现"""
from packages.domain import AssetLibrary, AssetLibraryKind
@@ -1,4 +1,5 @@
"""Asset InMemory Repository 实现"""
from packages.domain import Asset
@@ -1,4 +1,5 @@
"""项目管理 In-Memory Repository 实现"""
from packages.domain import Milestone, Task, TaskIssue
from packages.ports.project_management_repositories import (
MilestoneRepository,
+13 -11
View File
@@ -1,21 +1,23 @@
"""
用户仓储 In-Memory 实现
"""
from typing import Optional, Dict
from typing import Dict, Optional
from packages.domain.entities import User
from packages.ports.user_repository import UserRepository
class InMemoryUserRepository(UserRepository):
"""用户仓储内存实现"""
def __init__(self):
self._users: Dict[str, User] = {}
self._email_index: Dict[str, str] = {} # email -> user_id
self._username_index: Dict[str, str] = {} # username -> user_id
self._verification_token_index: Dict[str, str] = {} # token -> user_id
self._reset_token_index: Dict[str, str] = {} # token -> user_id
def save(self, user: User) -> None:
"""保存用户"""
self._users[user.id] = user
@@ -26,45 +28,45 @@ class InMemoryUserRepository(UserRepository):
self._verification_token_index[user.email_verification_token] = user.id
if user.password_reset_token:
self._reset_token_index[user.password_reset_token] = user.id
def find_by_id(self, user_id: str) -> Optional[User]:
"""根据 ID 查找用户"""
return self._users.get(user_id)
def find_by_email(self, email: str) -> Optional[User]:
"""根据邮箱查找用户"""
user_id = self._email_index.get(email.lower())
if user_id:
return self._users.get(user_id)
return None
def find_by_username(self, username: str) -> Optional[User]:
"""根据用户名查找用户"""
user_id = self._username_index.get(username.lower())
if user_id:
return self._users.get(user_id)
return None
def find_by_verification_token(self, token: str) -> Optional[User]:
"""根据邮箱验证令牌查找用户"""
user_id = self._verification_token_index.get(token)
if user_id:
return self._users.get(user_id)
return None
def find_by_password_reset_token(self, token: str) -> Optional[User]:
"""根据密码重置令牌查找用户"""
user_id = self._reset_token_index.get(token)
if user_id:
return self._users.get(user_id)
return None
def delete(self, user_id: str) -> bool:
"""删除用户"""
user = self._users.get(user_id)
if not user:
return False
# 清理索引
self._email_index.pop(user.email.lower(), None)
if user.username:
@@ -73,7 +75,7 @@ class InMemoryUserRepository(UserRepository):
self._verification_token_index.pop(user.email_verification_token, None)
if user.password_reset_token:
self._reset_token_index.pop(user.password_reset_token, None)
# 删除用户
del self._users[user_id]
return True
@@ -1,24 +1,26 @@
"""
WorkspaceInvitation 仓储 In-Memory 实现
"""
from typing import Optional, Dict
from packages.domain.entities import WorkspaceInvitation, InvitationStatus
from typing import Dict, Optional
from packages.domain.entities import InvitationStatus, WorkspaceInvitation
from packages.ports.workspace_invitation_repository import WorkspaceInvitationRepository
class InMemoryWorkspaceInvitationRepository(WorkspaceInvitationRepository):
"""WorkspaceInvitation 仓储内存实现"""
def __init__(self):
self._invitations: Dict[str, WorkspaceInvitation] = {}
self._token_index: Dict[str, str] = {} # token -> invitation_id
self._workspace_email_index: Dict[tuple[str, str], str] = {} # (workspace_id, email) -> invitation_id
def save(self, invitation: WorkspaceInvitation) -> None:
"""保存邀请"""
self._invitations[invitation.id] = invitation
self._token_index[invitation.invitation_token] = invitation.id
# 只为 pending 状态的邀请建立索引
if invitation.status == InvitationStatus.PENDING:
key = (invitation.workspace_id, invitation.invitee_email.lower())
@@ -27,18 +29,18 @@ class InMemoryWorkspaceInvitationRepository(WorkspaceInvitationRepository):
# 如果状态改变,清理索引
key = (invitation.workspace_id, invitation.invitee_email.lower())
self._workspace_email_index.pop(key, None)
def find_by_id(self, invitation_id: str) -> Optional[WorkspaceInvitation]:
"""根据 ID 查找邀请"""
return self._invitations.get(invitation_id)
def find_by_token(self, token: str) -> Optional[WorkspaceInvitation]:
"""根据令牌查找邀请"""
invitation_id = self._token_index.get(token)
if invitation_id:
return self._invitations.get(invitation_id)
return None
def find_pending_by_workspace_and_email(
self,
workspace_id: str,
@@ -50,18 +52,18 @@ class InMemoryWorkspaceInvitationRepository(WorkspaceInvitationRepository):
if invitation_id:
return self._invitations.get(invitation_id)
return None
def delete(self, invitation_id: str) -> bool:
"""删除邀请"""
invitation = self._invitations.get(invitation_id)
if not invitation:
return False
# 清理索引
self._token_index.pop(invitation.invitation_token, None)
key = (invitation.workspace_id, invitation.invitee_email.lower())
self._workspace_email_index.pop(key, None)
# 删除邀请
del self._invitations[invitation_id]
return True
@@ -1,42 +1,44 @@
"""
WorkspaceMember 仓储 In-Memory 实现
"""
from typing import Optional, Dict, List
from typing import Dict, List, Optional
from packages.domain.entities import WorkspaceMember
from packages.ports.workspace_member_repository import WorkspaceMemberRepository
class InMemoryWorkspaceMemberRepository(WorkspaceMemberRepository):
"""WorkspaceMember 仓储内存实现"""
def __init__(self):
self._members: Dict[str, WorkspaceMember] = {}
self._workspace_user_index: Dict[tuple[str, str], str] = {} # (workspace_id, user_id) -> member_id
self._user_index: Dict[str, List[str]] = {} # user_id -> [member_ids]
self._workspace_index: Dict[str, List[str]] = {} # workspace_id -> [member_ids]
def save(self, member: WorkspaceMember) -> None:
"""保存成员"""
self._members[member.id] = member
# 更新索引
key = (member.workspace_id, member.user_id)
self._workspace_user_index[key] = member.id
if member.user_id not in self._user_index:
self._user_index[member.user_id] = []
if member.id not in self._user_index[member.user_id]:
self._user_index[member.user_id].append(member.id)
if member.workspace_id not in self._workspace_index:
self._workspace_index[member.workspace_id] = []
if member.id not in self._workspace_index[member.workspace_id]:
self._workspace_index[member.workspace_id].append(member.id)
def find_by_id(self, member_id: str) -> Optional[WorkspaceMember]:
"""根据 ID 查找成员"""
return self._members.get(member_id)
def find_by_workspace_and_user(
self,
workspace_id: str,
@@ -48,37 +50,37 @@ class InMemoryWorkspaceMemberRepository(WorkspaceMemberRepository):
if member_id:
return self._members.get(member_id)
return None
def find_by_user(self, user_id: str) -> List[WorkspaceMember]:
"""查找用户的所有成员记录"""
member_ids = self._user_index.get(user_id, [])
return [self._members[mid] for mid in member_ids if mid in self._members]
def find_by_workspace(self, workspace_id: str) -> List[WorkspaceMember]:
"""查找 workspace 的所有成员"""
member_ids = self._workspace_index.get(workspace_id, [])
return [self._members[mid] for mid in member_ids if mid in self._members]
def count_by_workspace(self, workspace_id: str) -> int:
"""统计 workspace 的成员数量"""
return len(self._workspace_index.get(workspace_id, []))
def delete(self, member_id: str) -> bool:
"""删除成员"""
member = self._members.get(member_id)
if not member:
return False
# 清理索引
key = (member.workspace_id, member.user_id)
self._workspace_user_index.pop(key, None)
if member.user_id in self._user_index:
self._user_index[member.user_id].remove(member_id)
if member.workspace_id in self._workspace_index:
self._workspace_index[member.workspace_id].remove(member_id)
# 删除成员
del self._members[member_id]
return True
@@ -1,25 +1,27 @@
"""
Workspace 仓储 In-Memory 实现
"""
from typing import Optional, Dict
from typing import Dict, Optional
from packages.domain.entities import Workspace
from packages.ports.workspace_repository import WorkspaceRepository
class InMemoryWorkspaceRepository(WorkspaceRepository):
"""Workspace 仓储内存实现"""
def __init__(self):
self._workspaces: Dict[str, Workspace] = {}
def save(self, workspace: Workspace) -> None:
"""保存 Workspace"""
self._workspaces[workspace.id] = workspace
def find_by_id(self, workspace_id: str) -> Optional[Workspace]:
"""根据 ID 查找 Workspace"""
return self._workspaces.get(workspace_id)
def delete(self, workspace_id: str) -> bool:
"""删除 Workspace"""
if workspace_id in self._workspaces:
+1 -3
View File
@@ -5,6 +5,4 @@ package is kept only as a migration marker; do not import it in application,
API, worker, or new tests.
"""
raise RuntimeError(
"packages.adapters.postgres is deprecated; use packages.adapters.sqlalchemy_impl instead"
)
raise RuntimeError("packages.adapters.postgres is deprecated; use packages.adapters.sqlalchemy_impl instead")
+19 -5
View File
@@ -1,6 +1,7 @@
"""
Asset PostgreSQL Repository 实现
"""
import json
from sqlalchemy import and_, func, select
@@ -24,7 +25,7 @@ class PostgresAssetRepository(AssetRepository):
project_id=asset.project_id,
asset_library_id=asset.library_id,
name=asset.name,
file_type=asset.mime_type.split("/")[0] if "/" in asset.mime_type else asset.mime_type,
file_type=(asset.mime_type.split("/")[0] if "/" in asset.mime_type else asset.mime_type),
file_size=asset.file_size,
file_url=asset.storage_key,
thumbnail_url=asset.thumbnail_url,
@@ -35,7 +36,7 @@ class PostgresAssetRepository(AssetRepository):
codec=asset.codec,
status=asset.status.value,
classification_status=asset.classification_status.value,
classification_result=json.dumps(asset.metadata) if asset.metadata else None,
classification_result=(json.dumps(asset.metadata) if asset.metadata else None),
quality_score=asset.quality_score,
uploaded_by_user_id=asset.uploaded_by_user_id or "system",
created_at=asset.created_at,
@@ -59,7 +60,12 @@ class PostgresAssetRepository(AssetRepository):
) -> list[Asset]:
result = await self.session.execute(
select(AssetModel)
.where(and_(AssetModel.project_id == project_id, AssetModel.workspace_id == workspace_id))
.where(
and_(
AssetModel.project_id == project_id,
AssetModel.workspace_id == workspace_id,
)
)
.order_by(AssetModel.created_at.desc())
.offset(skip)
.limit(limit)
@@ -75,7 +81,12 @@ class PostgresAssetRepository(AssetRepository):
) -> list[Asset]:
result = await self.session.execute(
select(AssetModel)
.where(and_(AssetModel.asset_library_id == library_id, AssetModel.workspace_id == workspace_id))
.where(
and_(
AssetModel.asset_library_id == library_id,
AssetModel.workspace_id == workspace_id,
)
)
.order_by(AssetModel.created_at.desc())
.offset(skip)
.limit(limit)
@@ -118,7 +129,10 @@ class PostgresAssetRepository(AssetRepository):
async def count_by_project(self, project_id: str, workspace_id: str) -> int:
result = await self.session.execute(
select(func.count(AssetModel.id)).where(
and_(AssetModel.project_id == project_id, AssetModel.workspace_id == workspace_id)
and_(
AssetModel.project_id == project_id,
AssetModel.workspace_id == workspace_id,
)
)
)
return result.scalar() or 0
+12 -10
View File
@@ -1,7 +1,9 @@
"""
数据库连接池管理
"""
from typing import Optional
import psycopg2
from psycopg2 import pool
from psycopg2.extras import RealDictCursor
@@ -9,15 +11,15 @@ from psycopg2.extras import RealDictCursor
class DatabaseConnectionPool:
"""PostgreSQL 连接池"""
_instance: Optional['DatabaseConnectionPool'] = None
_instance: Optional["DatabaseConnectionPool"] = None
_pool: Optional[pool.ThreadedConnectionPool] = None
def __new__(cls):
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def initialize(
self,
connection_string: str,
@@ -31,18 +33,18 @@ class DatabaseConnectionPool:
maxconn=maxconn,
dsn=connection_string,
)
def get_connection(self):
"""从连接池获取连接"""
if self._pool is None:
raise RuntimeError("Connection pool not initialized")
return self._pool.getconn()
def put_connection(self, conn):
"""将连接归还到连接池"""
if self._pool is not None:
self._pool.putconn(conn)
def close_all(self):
"""关闭所有连接"""
if self._pool is not None:
@@ -56,17 +58,17 @@ db_pool = DatabaseConnectionPool()
class PooledConnection:
"""连接池上下文管理器"""
def __init__(self, cursor_factory=RealDictCursor):
self.cursor_factory = cursor_factory
self.conn = None
def __enter__(self):
self.conn = db_pool.get_connection()
if self.cursor_factory:
self.conn.cursor_factory = self.cursor_factory
return self.conn
def __exit__(self, exc_type, exc_val, exc_tb):
if self.conn:
if exc_type is not None:
@@ -1,7 +1,9 @@
"""
PostgreSQL Project Repository 实现
"""
from typing import Optional, List
from typing import List, Optional
import psycopg2
from psycopg2.extras import RealDictCursor
@@ -11,21 +13,23 @@ from packages.ports.project_repository import ProjectRepository
class PostgresProjectRepository(ProjectRepository):
"""Project 仓储 PostgreSQL 实现"""
def __init__(self, connection_string: str):
self.connection_string = connection_string
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, project: Project) -> None:
"""保存项目"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("""
cur.execute(
"""
INSERT INTO projects (
id, workspace_id, name, description, status,
created_by, created_at, updated_at
@@ -38,20 +42,22 @@ class PostgresProjectRepository(ProjectRepository):
description = EXCLUDED.description,
status = EXCLUDED.status,
updated_at = EXCLUDED.updated_at
""", {
"id": project.id,
"workspace_id": project.workspace_id,
"name": project.name,
"description": project.description,
"status": project.status,
"created_by": project.created_by,
"created_at": project.created_at,
"updated_at": project.updated_at,
})
""",
{
"id": project.id,
"workspace_id": project.workspace_id,
"name": project.name,
"description": project.description,
"status": project.status,
"created_by": project.created_by,
"created_at": project.created_at,
"updated_at": project.updated_at,
},
)
conn.commit()
finally:
conn.close()
def find_by_id(self, project_id: str) -> Optional[Project]:
"""根据 ID 查找项目"""
conn = self._get_connection()
@@ -62,7 +68,7 @@ class PostgresProjectRepository(ProjectRepository):
return self._row_to_project(row) if row else None
finally:
conn.close()
def find_by_workspace(self, workspace_id: str) -> List[Project]:
"""根据 workspace 查找所有项目"""
conn = self._get_connection()
@@ -70,13 +76,13 @@ class PostgresProjectRepository(ProjectRepository):
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM projects WHERE workspace_id = %s ORDER BY created_at DESC",
(workspace_id,)
(workspace_id,),
)
rows = cur.fetchall()
return [self._row_to_project(row) for row in rows]
finally:
conn.close()
def find_by_creator(self, user_id: str) -> List[Project]:
"""根据创建者查找项目"""
conn = self._get_connection()
@@ -84,13 +90,13 @@ class PostgresProjectRepository(ProjectRepository):
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM projects WHERE created_by = %s ORDER BY created_at DESC",
(user_id,)
(user_id,),
)
rows = cur.fetchall()
return [self._row_to_project(row) for row in rows]
finally:
conn.close()
def count_by_workspace(self, workspace_id: str) -> int:
"""统计 workspace 的项目数量"""
conn = self._get_connection()
@@ -98,12 +104,12 @@ class PostgresProjectRepository(ProjectRepository):
with conn.cursor() as cur:
cur.execute(
"SELECT COUNT(*) FROM projects WHERE workspace_id = %s",
(workspace_id,)
(workspace_id,),
)
return cur.fetchone()["count"]
finally:
conn.close()
def delete(self, project_id: str) -> bool:
"""删除项目"""
conn = self._get_connection()
@@ -115,7 +121,7 @@ class PostgresProjectRepository(ProjectRepository):
return deleted
finally:
conn.close()
def _row_to_project(self, row: dict) -> Project:
"""将数据库行转换为 Project 对象"""
return Project(
+37 -31
View File
@@ -1,10 +1,12 @@
"""
PostgreSQL User Repository 实现
"""
from datetime import datetime
from typing import Optional
import psycopg2
from psycopg2.extras import RealDictCursor
from datetime import datetime
from packages.domain.entities import User
from packages.ports.user_repository import UserRepository
@@ -12,22 +14,24 @@ from packages.ports.user_repository import UserRepository
class PostgresUserRepository(UserRepository):
"""User 仓储 PostgreSQL 实现"""
def __init__(self, connection_string: str):
self.connection_string = connection_string
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, user: User) -> None:
"""保存用户"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
# Upsert (插入或更新)
cur.execute("""
cur.execute(
"""
INSERT INTO users (
id, email, display_name, username, password_hash,
email_verified, email_verification_token,
@@ -50,24 +54,26 @@ class PostgresUserRepository(UserRepository):
password_reset_expires_at = EXCLUDED.password_reset_expires_at,
last_login_at = EXCLUDED.last_login_at,
last_login_ip = EXCLUDED.last_login_ip
""", {
"id": user.id,
"email": user.email,
"display_name": user.display_name,
"username": user.username,
"password_hash": user.password_hash,
"email_verified": user.email_verified,
"email_verification_token": user.email_verification_token,
"password_reset_token": user.password_reset_token,
"password_reset_expires_at": user.password_reset_expires_at,
"last_login_at": user.last_login_at,
"last_login_ip": user.last_login_ip,
"created_at": user.created_at,
})
""",
{
"id": user.id,
"email": user.email,
"display_name": user.display_name,
"username": user.username,
"password_hash": user.password_hash,
"email_verified": user.email_verified,
"email_verification_token": user.email_verification_token,
"password_reset_token": user.password_reset_token,
"password_reset_expires_at": user.password_reset_expires_at,
"last_login_at": user.last_login_at,
"last_login_ip": user.last_login_ip,
"created_at": user.created_at,
},
)
conn.commit()
finally:
conn.close()
def find_by_id(self, user_id: str) -> Optional[User]:
"""根据 ID 查找用户"""
conn = self._get_connection()
@@ -75,13 +81,13 @@ class PostgresUserRepository(UserRepository):
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE id = %s", (user_id,))
row = cur.fetchone()
if row:
return self._row_to_user(row)
return None
finally:
conn.close()
def find_by_email(self, email: str) -> Optional[User]:
"""根据邮箱查找用户"""
conn = self._get_connection()
@@ -89,13 +95,13 @@ class PostgresUserRepository(UserRepository):
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE email = %s", (email.lower(),))
row = cur.fetchone()
if row:
return self._row_to_user(row)
return None
finally:
conn.close()
def find_by_username(self, username: str) -> Optional[User]:
"""根据用户名查找用户"""
conn = self._get_connection()
@@ -103,13 +109,13 @@ class PostgresUserRepository(UserRepository):
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE username = %s", (username.lower(),))
row = cur.fetchone()
if row:
return self._row_to_user(row)
return None
finally:
conn.close()
def find_by_verification_token(self, token: str) -> Optional[User]:
"""根据邮箱验证令牌查找用户"""
conn = self._get_connection()
@@ -117,13 +123,13 @@ class PostgresUserRepository(UserRepository):
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE email_verification_token = %s", (token,))
row = cur.fetchone()
if row:
return self._row_to_user(row)
return None
finally:
conn.close()
def find_by_password_reset_token(self, token: str) -> Optional[User]:
"""根据密码重置令牌查找用户"""
conn = self._get_connection()
@@ -131,13 +137,13 @@ class PostgresUserRepository(UserRepository):
with conn.cursor() as cur:
cur.execute("SELECT * FROM users WHERE password_reset_token = %s", (token,))
row = cur.fetchone()
if row:
return self._row_to_user(row)
return None
finally:
conn.close()
def delete(self, user_id: str) -> bool:
"""删除用户"""
conn = self._get_connection()
@@ -149,7 +155,7 @@ class PostgresUserRepository(UserRepository):
return deleted
finally:
conn.close()
def _row_to_user(self, row: dict) -> User:
"""将数据库行转换为 User 对象"""
return User(
@@ -1,7 +1,9 @@
"""
PostgreSQL WorkspaceInvitation Repository 实现
"""
from typing import Optional, List
from typing import List, Optional
import psycopg2
from psycopg2.extras import RealDictCursor
@@ -11,21 +13,23 @@ from packages.ports.workspace_invitation_repository import WorkspaceInvitationRe
class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
"""WorkspaceInvitation 仓储 PostgreSQL 实现"""
def __init__(self, connection_string: str):
self.connection_string = connection_string
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, invitation: WorkspaceInvitation) -> None:
"""保存邀请"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("""
cur.execute(
"""
INSERT INTO workspace_invitations (
id, workspace_id, email, role, token,
invited_by, expires_at, status, created_at
@@ -35,32 +39,37 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
)
ON CONFLICT (id) DO UPDATE SET
status = EXCLUDED.status
""", {
"id": invitation.id,
"workspace_id": invitation.workspace_id,
"email": invitation.email,
"role": invitation.role,
"token": invitation.token,
"invited_by": invitation.invited_by,
"expires_at": invitation.expires_at,
"status": invitation.status,
"created_at": invitation.created_at,
})
""",
{
"id": invitation.id,
"workspace_id": invitation.workspace_id,
"email": invitation.email,
"role": invitation.role,
"token": invitation.token,
"invited_by": invitation.invited_by,
"expires_at": invitation.expires_at,
"status": invitation.status,
"created_at": invitation.created_at,
},
)
conn.commit()
finally:
conn.close()
def find_by_id(self, invitation_id: str) -> Optional[WorkspaceInvitation]:
"""根据 ID 查找邀请"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("SELECT * FROM workspace_invitations WHERE id = %s", (invitation_id,))
cur.execute(
"SELECT * FROM workspace_invitations WHERE id = %s",
(invitation_id,),
)
row = cur.fetchone()
return self._row_to_invitation(row) if row else None
finally:
conn.close()
def find_by_token(self, token: str) -> Optional[WorkspaceInvitation]:
"""根据 token 查找邀请"""
conn = self._get_connection()
@@ -71,7 +80,7 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
return self._row_to_invitation(row) if row else None
finally:
conn.close()
def find_by_email(self, email: str) -> List[WorkspaceInvitation]:
"""根据邮箱查找所有邀请"""
conn = self._get_connection()
@@ -79,28 +88,31 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM workspace_invitations WHERE email = %s ORDER BY created_at DESC",
(email,)
(email,),
)
rows = cur.fetchall()
return [self._row_to_invitation(row) for row in rows]
finally:
conn.close()
def find_pending_by_email(self, email: str) -> List[WorkspaceInvitation]:
"""查找邮箱的待处理邀请"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("""
cur.execute(
"""
SELECT * FROM workspace_invitations
WHERE email = %s AND status = 'pending' AND expires_at > NOW()
ORDER BY created_at DESC
""", (email,))
""",
(email,),
)
rows = cur.fetchall()
return [self._row_to_invitation(row) for row in rows]
finally:
conn.close()
def delete(self, invitation_id: str) -> bool:
"""删除邀请"""
conn = self._get_connection()
@@ -112,7 +124,7 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
return deleted
finally:
conn.close()
def _row_to_invitation(self, row: dict) -> WorkspaceInvitation:
"""将数据库行转换为 WorkspaceInvitation 对象"""
return WorkspaceInvitation(
@@ -1,7 +1,9 @@
"""
PostgreSQL WorkspaceMember Repository 实现
"""
from typing import Optional, List
from typing import List, Optional
import psycopg2
from psycopg2.extras import RealDictCursor
@@ -11,21 +13,23 @@ from packages.ports.workspace_member_repository import WorkspaceMemberRepository
class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
"""WorkspaceMember 仓储 PostgreSQL 实现"""
def __init__(self, connection_string: str):
self.connection_string = connection_string
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, member: WorkspaceMember) -> None:
"""保存成员"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("""
cur.execute(
"""
INSERT INTO workspace_members (
id, workspace_id, user_id, role, invited_by, joined_at
) VALUES (
@@ -34,18 +38,20 @@ class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
)
ON CONFLICT (workspace_id, user_id) DO UPDATE SET
role = EXCLUDED.role
""", {
"id": member.id,
"workspace_id": member.workspace_id,
"user_id": member.user_id,
"role": member.role,
"invited_by": member.invited_by,
"joined_at": member.joined_at,
})
""",
{
"id": member.id,
"workspace_id": member.workspace_id,
"user_id": member.user_id,
"role": member.role,
"invited_by": member.invited_by,
"joined_at": member.joined_at,
},
)
conn.commit()
finally:
conn.close()
def find_by_id(self, member_id: str) -> Optional[WorkspaceMember]:
"""根据 ID 查找成员"""
conn = self._get_connection()
@@ -56,7 +62,7 @@ class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
return self._row_to_member(row) if row else None
finally:
conn.close()
def find_by_workspace_and_user(
self,
workspace_id: str,
@@ -68,13 +74,13 @@ class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM workspace_members WHERE workspace_id = %s AND user_id = %s",
(workspace_id, user_id)
(workspace_id, user_id),
)
row = cur.fetchone()
return self._row_to_member(row) if row else None
finally:
conn.close()
def find_by_user(self, user_id: str) -> List[WorkspaceMember]:
"""查找用户的所有成员记录"""
conn = self._get_connection()
@@ -82,13 +88,13 @@ class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM workspace_members WHERE user_id = %s ORDER BY joined_at DESC",
(user_id,)
(user_id,),
)
rows = cur.fetchall()
return [self._row_to_member(row) for row in rows]
finally:
conn.close()
def find_by_workspace(self, workspace_id: str) -> List[WorkspaceMember]:
"""查找 workspace 的所有成员"""
conn = self._get_connection()
@@ -96,13 +102,13 @@ class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
with conn.cursor() as cur:
cur.execute(
"SELECT * FROM workspace_members WHERE workspace_id = %s ORDER BY joined_at",
(workspace_id,)
(workspace_id,),
)
rows = cur.fetchall()
return [self._row_to_member(row) for row in rows]
finally:
conn.close()
def count_by_workspace(self, workspace_id: str) -> int:
"""统计 workspace 的成员数量"""
conn = self._get_connection()
@@ -110,12 +116,12 @@ class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
with conn.cursor() as cur:
cur.execute(
"SELECT COUNT(*) FROM workspace_members WHERE workspace_id = %s",
(workspace_id,)
(workspace_id,),
)
return cur.fetchone()["count"]
finally:
conn.close()
def delete(self, member_id: str) -> bool:
"""删除成员"""
conn = self._get_connection()
@@ -127,7 +133,7 @@ class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
return deleted
finally:
conn.close()
def _row_to_member(self, row: dict) -> WorkspaceMember:
"""将数据库行转换为 WorkspaceMember 对象"""
return WorkspaceMember(
@@ -1,7 +1,9 @@
"""
PostgreSQL Workspace Repository 实现
"""
from typing import Optional
import psycopg2
from psycopg2.extras import RealDictCursor
@@ -11,21 +13,23 @@ from packages.ports.workspace_repository import WorkspaceRepository
class PostgresWorkspaceRepository(WorkspaceRepository):
"""Workspace 仓储 PostgreSQL 实现"""
def __init__(self, connection_string: str):
self.connection_string = connection_string
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, workspace: Workspace) -> None:
"""保存工作空间"""
conn = self._get_connection()
try:
with conn.cursor() as cur:
cur.execute("""
cur.execute(
"""
INSERT INTO workspaces (
id, name, owner_user_id, subscription_plan,
subscription_status, subscription_expires_at,
@@ -43,22 +47,24 @@ class PostgresWorkspaceRepository(WorkspaceRepository):
max_projects = EXCLUDED.max_projects,
max_storage_gb = EXCLUDED.max_storage_gb,
used_storage_gb = EXCLUDED.used_storage_gb
""", {
"id": workspace.id,
"name": workspace.name,
"owner_user_id": workspace.owner_user_id,
"subscription_plan": workspace.subscription_plan,
"subscription_status": workspace.subscription_status,
"subscription_expires_at": workspace.subscription_expires_at,
"max_projects": workspace.max_projects,
"max_storage_gb": workspace.max_storage_gb,
"used_storage_gb": workspace.used_storage_gb,
"created_at": workspace.created_at,
})
""",
{
"id": workspace.id,
"name": workspace.name,
"owner_user_id": workspace.owner_user_id,
"subscription_plan": workspace.subscription_plan,
"subscription_status": workspace.subscription_status,
"subscription_expires_at": workspace.subscription_expires_at,
"max_projects": workspace.max_projects,
"max_storage_gb": workspace.max_storage_gb,
"used_storage_gb": workspace.used_storage_gb,
"created_at": workspace.created_at,
},
)
conn.commit()
finally:
conn.close()
def find_by_id(self, workspace_id: str) -> Optional[Workspace]:
"""根据 ID 查找工作空间"""
conn = self._get_connection()
@@ -66,13 +72,13 @@ class PostgresWorkspaceRepository(WorkspaceRepository):
with conn.cursor() as cur:
cur.execute("SELECT * FROM workspaces WHERE id = %s", (workspace_id,))
row = cur.fetchone()
if row:
return self._row_to_workspace(row)
return None
finally:
conn.close()
def delete(self, workspace_id: str) -> bool:
"""删除工作空间"""
conn = self._get_connection()
@@ -84,7 +90,7 @@ class PostgresWorkspaceRepository(WorkspaceRepository):
return deleted
finally:
conn.close()
def _row_to_workspace(self, row: dict) -> Workspace:
"""将数据库行转换为 Workspace 对象"""
return Workspace(
+5 -1
View File
@@ -1,3 +1,7 @@
from packages.adapters.redis.session_store import RedisConfig, SessionStore, get_session_store
from packages.adapters.redis.session_store import (
RedisConfig,
SessionStore,
get_session_store,
)
__all__ = ["RedisConfig", "SessionStore", "get_session_store"]
+57 -66
View File
@@ -2,15 +2,18 @@
Redis Session 存储
用于存储 refresh_token 和 Session 信息
"""
from typing import Optional
from datetime import datetime, timedelta, timezone
import json
from datetime import datetime, timedelta, timezone
from typing import Optional
import redis
from redis import Redis
class RedisConfig:
"""Redis 配置"""
HOST: str = "localhost"
PORT: int = 6379
DB: int = 0
@@ -20,11 +23,11 @@ class RedisConfig:
class SessionStore:
"""Session 存储服务"""
def __init__(self, redis_client: Optional[Redis] = None, config: Optional[RedisConfig] = None):
"""
初始化 Session 存储
Args:
redis_client: Redis 客户端(可选,用于注入)
config: Redis 配置(可选)
@@ -40,19 +43,19 @@ class SessionStore:
password=cfg.PASSWORD,
decode_responses=cfg.DECODE_RESPONSES,
)
def _session_key(self, session_id: str) -> str:
"""生成 Session key"""
return f"session:{session_id}"
def _refresh_token_key(self, session_id: str) -> str:
"""生成 refresh_token key"""
return f"refresh_token:{session_id}"
def _user_sessions_key(self, user_id: str) -> str:
"""生成用户所有 Session 的 key"""
return f"user_sessions:{user_id}"
def save_session(
self,
session_id: str,
@@ -64,7 +67,7 @@ class SessionStore:
) -> bool:
"""
保存 Session
Args:
session_id: Session ID
user_id: 用户 ID
@@ -72,14 +75,14 @@ class SessionStore:
device_info: 设备信息
ip_address: IP 地址
expires_in_seconds: 过期时间(秒)
Returns:
是否保存成功
"""
try:
now = datetime.now(timezone.utc)
expires_at = now + timedelta(seconds=expires_in_seconds)
session_data = {
"session_id": session_id,
"user_id": user_id,
@@ -89,61 +92,53 @@ class SessionStore:
"last_active_at": now.isoformat(),
"expires_at": expires_at.isoformat(),
}
# 保存 Session 数据
session_key = self._session_key(session_id)
self.redis.setex(
session_key,
expires_in_seconds,
json.dumps(session_data)
)
self.redis.setex(session_key, expires_in_seconds, json.dumps(session_data))
# 保存 refresh_token 映射
refresh_token_key = self._refresh_token_key(session_id)
self.redis.setex(
refresh_token_key,
expires_in_seconds,
refresh_token
)
self.redis.setex(refresh_token_key, expires_in_seconds, refresh_token)
# 添加到用户的 Session 集合
user_sessions_key = self._user_sessions_key(user_id)
self.redis.sadd(user_sessions_key, session_id)
self.redis.expire(user_sessions_key, expires_in_seconds)
return True
except Exception as e:
print(f"Failed to save session: {e}")
return False
def get_session(self, session_id: str) -> Optional[dict]:
"""
获取 Session
Args:
session_id: Session ID
Returns:
Session 数据,如果不存在返回 None
"""
try:
session_key = self._session_key(session_id)
data = self.redis.get(session_key)
if data:
return json.loads(data)
return None
except Exception as e:
print(f"Failed to get session: {e}")
return None
def get_refresh_token(self, session_id: str) -> Optional[str]:
"""
获取 refresh_token
Args:
session_id: Session ID
Returns:
refresh_token,如果不存在返回 None
"""
@@ -153,14 +148,14 @@ class SessionStore:
except Exception as e:
print(f"Failed to get refresh_token: {e}")
return None
def update_last_active(self, session_id: str) -> bool:
"""
更新 Session 最后活跃时间
Args:
session_id: Session ID
Returns:
是否更新成功
"""
@@ -168,32 +163,28 @@ class SessionStore:
session = self.get_session(session_id)
if not session:
return False
session["last_active_at"] = datetime.now(timezone.utc).isoformat()
session_key = self._session_key(session_id)
ttl = self.redis.ttl(session_key)
if ttl > 0:
self.redis.setex(
session_key,
ttl,
json.dumps(session)
)
self.redis.setex(session_key, ttl, json.dumps(session))
return True
return False
except Exception as e:
print(f"Failed to update last active: {e}")
return False
def delete_session(self, session_id: str) -> bool:
"""
删除 Session(登出)
Args:
session_id: Session ID
Returns:
是否删除成功
"""
@@ -201,85 +192,85 @@ class SessionStore:
session = self.get_session(session_id)
if not session:
return False
user_id = session["user_id"]
# 删除 Session 数据
session_key = self._session_key(session_id)
self.redis.delete(session_key)
# 删除 refresh_token
refresh_token_key = self._refresh_token_key(session_id)
self.redis.delete(refresh_token_key)
# 从用户 Session 集合中移除
user_sessions_key = self._user_sessions_key(user_id)
self.redis.srem(user_sessions_key, session_id)
return True
except Exception as e:
print(f"Failed to delete session: {e}")
return False
def get_user_sessions(self, user_id: str) -> list[dict]:
"""
获取用户的所有活跃 Session
Args:
user_id: 用户 ID
Returns:
Session 列表
"""
try:
user_sessions_key = self._user_sessions_key(user_id)
session_ids = self.redis.smembers(user_sessions_key)
sessions = []
for session_id in session_ids:
session = self.get_session(session_id)
if session:
sessions.append(session)
return sessions
except Exception as e:
print(f"Failed to get user sessions: {e}")
return []
def delete_all_user_sessions(self, user_id: str) -> int:
"""
删除用户的所有 Session(强制登出所有设备)
Args:
user_id: 用户 ID
Returns:
删除的 Session 数量
"""
try:
sessions = self.get_user_sessions(user_id)
count = 0
for session in sessions:
if self.delete_session(session["session_id"]):
count += 1
# 清空用户 Session 集合
user_sessions_key = self._user_sessions_key(user_id)
self.redis.delete(user_sessions_key)
return count
except Exception as e:
print(f"Failed to delete all user sessions: {e}")
return 0
def session_exists(self, session_id: str) -> bool:
"""
检查 Session 是否存在
Args:
session_id: Session ID
Returns:
是否存在
"""
+5 -1
View File
@@ -1,3 +1,7 @@
from packages.adapters.smtp.email_service import EmailConfig, EmailService, get_email_service
from packages.adapters.smtp.email_service import (
EmailConfig,
EmailService,
get_email_service,
)
__all__ = ["EmailConfig", "EmailService", "get_email_service"]
+39 -41
View File
@@ -2,16 +2,18 @@
邮件服务
支持 SMTP 发送邮件(验证/重置密码/邀请等)
"""
import smtplib
from email.mime.text import MIMEText
from email.mime.multipart import MIMEMultipart
from typing import Optional, List
from dataclasses import dataclass
from email.mime.multipart import MIMEMultipart
from email.mime.text import MIMEText
from typing import List, Optional
@dataclass
class EmailConfig:
"""邮件配置"""
smtp_host: str = "smtp.gmail.com"
smtp_port: int = 587
smtp_user: str = ""
@@ -23,16 +25,16 @@ class EmailConfig:
class EmailService:
"""邮件服务类"""
def __init__(self, config: Optional[EmailConfig] = None):
"""
初始化邮件服务
Args:
config: 邮件配置
"""
self.config = config or EmailConfig()
def send_email(
self,
to_email: str,
@@ -44,7 +46,7 @@ class EmailService:
) -> tuple[bool, Optional[str]]:
"""
发送邮件
Args:
to_email: 收件人邮箱
subject: 邮件主题
@@ -52,7 +54,7 @@ class EmailService:
text_body: 纯文本正文(可选,作为 HTML 的备用)
cc: 抄送列表
bcc: 密送列表
Returns:
(是否成功, 错误信息)
"""
@@ -62,46 +64,42 @@ class EmailService:
msg["Subject"] = subject
msg["From"] = f"{self.config.from_name} <{self.config.from_email}>"
msg["To"] = to_email
if cc:
msg["Cc"] = ", ".join(cc)
# 添加纯文本正文
if text_body:
part1 = MIMEText(text_body, "plain", "utf-8")
msg.attach(part1)
# 添加 HTML 正文
part2 = MIMEText(html_body, "html", "utf-8")
msg.attach(part2)
# 连接 SMTP 服务器
with smtplib.SMTP(self.config.smtp_host, self.config.smtp_port) as server:
if self.config.use_tls:
server.starttls()
# 登录
if self.config.smtp_user and self.config.smtp_password:
server.login(self.config.smtp_user, self.config.smtp_password)
# 发送
recipients = [to_email]
if cc:
recipients.extend(cc)
if bcc:
recipients.extend(bcc)
server.sendmail(
self.config.from_email,
recipients,
msg.as_string()
)
server.sendmail(self.config.from_email, recipients, msg.as_string())
return True, None
except Exception as e:
return False, str(e)
def send_verification_email(
self,
to_email: str,
@@ -110,17 +108,17 @@ class EmailService:
) -> tuple[bool, Optional[str]]:
"""
发送邮箱验证邮件
Args:
to_email: 收件人邮箱
username: 用户名
verification_url: 验证链接
Returns:
(是否成功, 错误信息)
"""
subject = "验证您的邮箱 - 小虾 SaaS"
html_body = f"""
<!DOCTYPE html>
<html>
@@ -154,7 +152,7 @@ class EmailService:
</body>
</html>
"""
text_body = f"""
欢迎加入小虾 SaaS!
@@ -168,9 +166,9 @@ class EmailService:
如果您没有注册小虾 SaaS,请忽略此邮件。
"""
return self.send_email(to_email, subject, html_body, text_body)
def send_password_reset_email(
self,
to_email: str,
@@ -179,17 +177,17 @@ class EmailService:
) -> tuple[bool, Optional[str]]:
"""
发送密码重置邮件
Args:
to_email: 收件人邮箱
username: 用户名
reset_url: 重置链接
Returns:
(是否成功, 错误信息)
"""
subject = "重置您的密码 - 小虾 SaaS"
html_body = f"""
<!DOCTYPE html>
<html>
@@ -223,7 +221,7 @@ class EmailService:
</body>
</html>
"""
text_body = f"""
重置密码请求
@@ -237,9 +235,9 @@ class EmailService:
如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。
"""
return self.send_email(to_email, subject, html_body, text_body)
def send_workspace_invitation_email(
self,
to_email: str,
@@ -250,19 +248,19 @@ class EmailService:
) -> tuple[bool, Optional[str]]:
"""
发送 Workspace 邀请邮件
Args:
to_email: 收件人邮箱
inviter_name: 邀请人姓名
workspace_name: 工作空间名称
role: 角色(Admin/Member/Viewer
invitation_url: 邀请链接
Returns:
(是否成功, 错误信息)
"""
subject = f"{inviter_name} 邀请您加入 {workspace_name} - 小虾 SaaS"
role_names = {
"owner": "所有者",
"admin": "管理员",
@@ -270,7 +268,7 @@ class EmailService:
"viewer": "查看者",
}
role_display = role_names.get(role.lower(), role)
html_body = f"""
<!DOCTYPE html>
<html>
@@ -307,7 +305,7 @@ class EmailService:
</body>
</html>
"""
text_body = f"""
工作空间邀请
@@ -321,7 +319,7 @@ class EmailService:
如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。
"""
return self.send_email(to_email, subject, html_body, text_body)
@@ -7,7 +7,13 @@ from .generated_video_repository import SQLAlchemyGeneratedVideoRepository
from .generation_task_repository import SQLAlchemyGenerationTaskRepository
from .ingest_job_repository import SQLAlchemyIngestJobRepository
from .project_repository import SQLAlchemyProjectRepository
from .session import Base, build_engine, build_session_factory, ensure_database_exists, initialize_database
from .session import (
Base,
build_engine,
build_session_factory,
ensure_database_exists,
initialize_database,
)
__all__ = [
"Base",
@@ -29,7 +29,7 @@ class SQLAlchemyAssetRepository:
project_id=asset.project_id,
asset_library_id=asset.library_id,
name=asset.name,
file_type=asset.mime_type.split('/')[0] if '/' in asset.mime_type else asset.mime_type,
file_type=(asset.mime_type.split("/")[0] if "/" in asset.mime_type else asset.mime_type),
file_size=asset.file_size,
file_url=asset.storage_key,
thumbnail_url=asset.thumbnail_url,
@@ -40,9 +40,9 @@ class SQLAlchemyAssetRepository:
codec=asset.codec,
status=asset.status.value,
classification_status=asset.classification_status.value,
classification_result=json.dumps(asset.metadata) if asset.metadata else None,
classification_result=(json.dumps(asset.metadata) if asset.metadata else None),
quality_score=asset.quality_score,
uploaded_by_user_id=asset.uploaded_by_user_id or 'system',
uploaded_by_user_id=asset.uploaded_by_user_id or "system",
created_at=asset.created_at,
updated_at=now,
)
@@ -80,11 +80,11 @@ class SQLAlchemyAssetRepository:
except Exception:
metadata = {}
mime_type = model.file_type
if '/' not in mime_type:
if "/" not in mime_type:
mime_type = {
'video': 'video/mp4',
'audio': 'audio/mpeg',
'image': 'image/jpeg',
"video": "video/mp4",
"audio": "audio/mpeg",
"image": "image/jpeg",
}.get(mime_type, mime_type)
return Asset(
id=model.id,
@@ -55,5 +55,9 @@ class SQLAlchemyGeneratedVideoRepository:
return [self.get(model.id) for model in models if self.get(model.id) is not None]
def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]:
models = self.session.query(GeneratedVideoModel).filter(GeneratedVideoModel.generation_task_id == generation_task_id).all()
models = (
self.session.query(GeneratedVideoModel)
.filter(GeneratedVideoModel.generation_task_id == generation_task_id)
.all()
)
return [self.get(model.id) for model in models if self.get(model.id) is not None]
+2 -1
View File
@@ -1,6 +1,7 @@
from datetime import datetime, timezone
from sqlalchemy import Boolean, Column, DateTime, Float, String, Text, create_engine
from sqlalchemy.orm import declarative_base
from datetime import datetime, timezone
Base = declarative_base()
@@ -1,5 +1,7 @@
"""项目管理 SQLAlchemy Repository 实现"""
import json
from sqlalchemy.orm import Session
from packages.domain import Milestone, Task, TaskIssue
@@ -8,6 +10,7 @@ from packages.ports.project_management_repositories import (
TaskIssueRepository,
TaskRepository,
)
from .models import MilestoneModel, TaskIssueModel, TaskModel
@@ -59,7 +62,7 @@ class SQLAlchemyTaskRepository(TaskRepository):
model = self._session.query(TaskModel).filter(TaskModel.id == task.id).first()
if not model:
raise ValueError(f"Task {task.id} not found")
model.name = task.name
model.description = task.description
model.status = task.status.value
@@ -73,7 +76,7 @@ class SQLAlchemyTaskRepository(TaskRepository):
model.actual_end_date = task.actual_end_date
model.tags_json = json.dumps(task.tags, ensure_ascii=False)
model.updated_at = task.updated_at
self._session.commit()
return task
@@ -83,6 +86,7 @@ class SQLAlchemyTaskRepository(TaskRepository):
def _model_to_entity(self, model: TaskModel) -> Task:
from packages.domain.project_management import TaskPriority, TaskStatus
return Task(
id=model.id,
project_id=model.project_id,
@@ -141,14 +145,14 @@ class SQLAlchemyMilestoneRepository(MilestoneRepository):
model = self._session.query(MilestoneModel).filter(MilestoneModel.id == milestone.id).first()
if not model:
raise ValueError(f"Milestone {milestone.id} not found")
model.name = milestone.name
model.description = milestone.description
model.target_date = milestone.target_date
model.completed = milestone.completed
model.completed_at = milestone.completed_at
model.updated_at = milestone.updated_at
self._session.commit()
return milestone
@@ -213,13 +217,13 @@ class SQLAlchemyTaskIssueRepository(TaskIssueRepository):
model = self._session.query(TaskIssueModel).filter(TaskIssueModel.id == issue.id).first()
if not model:
raise ValueError(f"TaskIssue {issue.id} not found")
model.title = issue.title
model.description = issue.description
model.resolved = issue.resolved
model.resolved_at = issue.resolved_at
model.updated_at = issue.updated_at
self._session.commit()
return issue
+4 -2
View File
@@ -6,7 +6,6 @@ from sqlalchemy.orm import sessionmaker
from packages.adapters.sqlalchemy_impl.models import Base
SCHEMA_INIT_LOCK_ID = 2026061501
SessionLocal = None
@@ -77,5 +76,8 @@ def initialize_database(engine) -> None:
Base.metadata.create_all(bind=connection)
connection.commit()
finally:
connection.execute(text("SELECT pg_advisory_unlock(:lock_id)"), {"lock_id": SCHEMA_INIT_LOCK_ID})
connection.execute(
text("SELECT pg_advisory_unlock(:lock_id)"),
{"lock_id": SCHEMA_INIT_LOCK_ID},
)
connection.commit()
+2 -1
View File
@@ -1,8 +1,9 @@
"""SQLite Tracker Adapter"""
from .project_management_repositories import (
SQLiteTaskRepository,
SQLiteMilestoneRepository,
SQLiteTaskIssueRepository,
SQLiteTaskRepository,
)
__all__ = [
@@ -1,94 +1,127 @@
"""SQLite 实现的项目管理 Repository"""
import sqlite3
from typing import List, Optional
from datetime import datetime
from packages.domain.project_management import Task, Milestone, TaskIssue, TaskStatus, TaskPriority
from typing import List, Optional
from packages.domain.project_management import (
Milestone,
Task,
TaskIssue,
TaskPriority,
TaskStatus,
)
DB_PATH = "tracker.db"
class SQLiteTaskRepository:
"""基于 SQLite 的任务仓储"""
def get_by_id(self, task_id: str) -> Optional[Task]:
conn = sqlite3.connect(DB_PATH)
conn.row_factory = sqlite3.Row
cursor = conn.cursor()
cursor.execute("SELECT * FROM tasks WHERE id = ?", (task_id,))
row = cursor.fetchone()
conn.close()
if not row:
return None
return Task(
id=str(row['id']),
name=row['name'],
description=row['description'] or "",
status=TaskStatus(row['status']) if row['status'] else TaskStatus.PENDING,
priority=TaskPriority(row['priority']) if row['priority'] else TaskPriority.MEDIUM,
id=str(row["id"]),
name=row["name"],
description=row["description"] or "",
status=TaskStatus(row["status"]) if row["status"] else TaskStatus.PENDING,
priority=(TaskPriority(row["priority"]) if row["priority"] else TaskPriority.MEDIUM),
progress=0, # tracker.db 没有 progress 字段
project_id=row['phase'] or "xiaoxia-saas",
project_id=row["phase"] or "xiaoxia-saas",
workspace_id="xiaoxia-workspace",
assignee_user_id=row['assigned_to'] or "",
created_at=datetime.fromisoformat(row['created_at']) if row['created_at'] else datetime.now(),
updated_at=datetime.now()
assignee_user_id=row["assigned_to"] or "",
created_at=(datetime.fromisoformat(row["created_at"]) if row["created_at"] else datetime.now()),
updated_at=datetime.now(),
)
def list_by_project(self, project_id: str, skip: int = 0, limit: int = 100) -> List[Task]:
conn = sqlite3.connect(DB_PATH)
conn.row_factory = sqlite3.Row
cursor = conn.cursor()
# 返回所有任务(忽略 project_id 过滤,因为 tracker.db 使用 phase
cursor.execute("""
cursor.execute(
"""
SELECT * FROM tasks
ORDER BY created_at DESC
LIMIT ? OFFSET ?
""", (limit, skip))
""",
(limit, skip),
)
rows = cursor.fetchall()
conn.close()
tasks = []
for row in rows:
tasks.append(Task(
id=str(row['id']),
name=row['name'],
description=row['description'] or "",
status=TaskStatus(row['status']) if row['status'] else TaskStatus.PENDING,
priority=TaskPriority(row['priority']) if row['priority'] else TaskPriority.MEDIUM,
progress=0,
project_id=row['phase'] or "xiaoxia-saas",
workspace_id="xiaoxia-workspace",
assignee_user_id=row['assigned_to'] or "",
created_at=datetime.fromisoformat(row['created_at']) if row['created_at'] else datetime.now(),
updated_at=datetime.now()
))
tasks.append(
Task(
id=str(row["id"]),
name=row["name"],
description=row["description"] or "",
status=(TaskStatus(row["status"]) if row["status"] else TaskStatus.PENDING),
priority=(TaskPriority(row["priority"]) if row["priority"] else TaskPriority.MEDIUM),
progress=0,
project_id=row["phase"] or "xiaoxia-saas",
workspace_id="xiaoxia-workspace",
assignee_user_id=row["assigned_to"] or "",
created_at=(datetime.fromisoformat(row["created_at"]) if row["created_at"] else datetime.now()),
updated_at=datetime.now(),
)
)
return tasks
def save(self, task: Task) -> Task:
conn = sqlite3.connect(DB_PATH)
cursor = conn.cursor()
if task.id and task.id.isdigit():
# 更新现有任务
cursor.execute("""
cursor.execute(
"""
UPDATE tasks
SET name = ?, description = ?, status = ?, priority = ?, assigned_to = ?
WHERE id = ?
""", (task.name, task.description, task.status.value, task.priority.value,
task.assignee_user_id, task.id))
""",
(
task.name,
task.description,
task.status.value,
task.priority.value,
task.assignee_user_id,
task.id,
),
)
else:
# 创建新任务
cursor.execute("""
cursor.execute(
"""
INSERT INTO tasks (name, description, status, phase, priority, assigned_to, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?)
""", (task.name, task.description, task.status.value, task.project_id,
task.priority.value, task.assignee_user_id, datetime.now().isoformat()))
""",
(
task.name,
task.description,
task.status.value,
task.project_id,
task.priority.value,
task.assignee_user_id,
datetime.now().isoformat(),
),
)
task.id = str(cursor.lastrowid)
conn.commit()
conn.close()
return task
@@ -96,48 +129,66 @@ class SQLiteTaskRepository:
class SQLiteMilestoneRepository:
"""基于 SQLite 的里程碑仓储"""
def list_by_project(self, project_id: str) -> List[Milestone]:
conn = sqlite3.connect(DB_PATH)
conn.row_factory = sqlite3.Row
cursor = conn.cursor()
cursor.execute("SELECT * FROM milestones ORDER BY start_date")
rows = cursor.fetchall()
conn.close()
milestones = []
for row in rows:
milestones.append(Milestone(
id=str(row['id']),
name=row['name'],
description=row['description'] or "",
target_date=row['end_date'] or "",
project_id=row['phase'] or "xiaoxia-saas",
workspace_id="xiaoxia-workspace",
created_at=datetime.fromisoformat(row['start_date']) if row['start_date'] else datetime.now()
))
milestones.append(
Milestone(
id=str(row["id"]),
name=row["name"],
description=row["description"] or "",
target_date=row["end_date"] or "",
project_id=row["phase"] or "xiaoxia-saas",
workspace_id="xiaoxia-workspace",
created_at=(datetime.fromisoformat(row["start_date"]) if row["start_date"] else datetime.now()),
)
)
return milestones
def save(self, milestone: Milestone) -> Milestone:
conn = sqlite3.connect(DB_PATH)
cursor = conn.cursor()
if milestone.id and milestone.id.isdigit():
cursor.execute("""
cursor.execute(
"""
UPDATE milestones
SET name = ?, description = ?, end_date = ?
WHERE id = ?
""", (milestone.name, milestone.description, milestone.target_date, milestone.id))
""",
(
milestone.name,
milestone.description,
milestone.target_date,
milestone.id,
),
)
else:
cursor.execute("""
cursor.execute(
"""
INSERT INTO milestones (name, description, phase, start_date, end_date)
VALUES (?, ?, ?, ?, ?)
""", (milestone.name, milestone.description, milestone.project_id,
datetime.now().isoformat(), milestone.target_date))
""",
(
milestone.name,
milestone.description,
milestone.project_id,
datetime.now().isoformat(),
milestone.target_date,
),
)
milestone.id = str(cursor.lastrowid)
conn.commit()
conn.close()
return milestone
@@ -145,9 +196,9 @@ class SQLiteMilestoneRepository:
class SQLiteTaskIssueRepository:
"""空实现 - tracker.db 没有 issues 表"""
def list_by_task(self, task_id: str) -> List[TaskIssue]:
return []
def save(self, issue: TaskIssue) -> TaskIssue:
return issue
+14 -3
View File
@@ -1,15 +1,26 @@
"""Application use cases package."""
from .asset_libraries import CreateAssetLibraryCommand, CreateAssetLibraryUseCase, ListAssetLibrariesUseCase
from .asset_libraries import (
CreateAssetLibraryCommand,
CreateAssetLibraryUseCase,
ListAssetLibrariesUseCase,
)
from .assets import CreateAssetCommand, CreateAssetUseCase, ListAssetsUseCase
from .classification_jobs import SubmitClassificationJobCommand, SubmitClassificationJobUseCase
from .classification_jobs import (
SubmitClassificationJobCommand,
SubmitClassificationJobUseCase,
)
from .generated_videos import (
GetGeneratedVideoDownloadUrlUseCase,
GetGeneratedVideoUseCase,
ListGeneratedVideosByTaskUseCase,
ListGeneratedVideosUseCase,
)
from .generation_tasks import CreateGenerationTaskCommand, CreateGenerationTaskUseCase, GetGenerationTaskUseCase
from .generation_tasks import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
GetGenerationTaskUseCase,
)
from .ingest_jobs import SubmitIngestJobCommand, SubmitIngestJobUseCase
from .projects import CreateProjectCommand, CreateProjectUseCase, ListProjectsUseCase
+14 -13
View File
@@ -1,25 +1,26 @@
"""认证相关 Use Cases"""
from packages.application.auth.register_user_use_case import (
RegisterUserUseCase,
RegisterUserRequest,
RegisterUserResponse,
VerifyEmailUseCase,
VerifyEmailRequest,
)
from packages.application.auth.login_use_case import (
LoginUseCase,
LoginRequest,
LoginResponse,
RefreshTokenUseCase,
RefreshTokenRequest,
LogoutUseCase,
LoginUseCase,
LogoutRequest,
LogoutUseCase,
RefreshTokenRequest,
RefreshTokenUseCase,
)
from packages.application.auth.password_reset_use_case import (
RequestPasswordResetUseCase,
RequestPasswordResetRequest,
ResetPasswordUseCase,
RequestPasswordResetUseCase,
ResetPasswordRequest,
ResetPasswordUseCase,
)
from packages.application.auth.register_user_use_case import (
RegisterUserRequest,
RegisterUserResponse,
RegisterUserUseCase,
VerifyEmailRequest,
VerifyEmailUseCase,
)
__all__ = [
+48 -46
View File
@@ -1,20 +1,18 @@
"""
用户登录 Use Case
"""
import secrets
from datetime import datetime, timezone, timedelta
from datetime import datetime, timedelta, timezone
from typing import Optional
from packages.adapters.redis import get_session_store
from packages.domain.auth import (
password_hasher,
jwt_service,
)
from packages.domain.auth import jwt_service, password_hasher
class LoginRequest:
"""登录请求"""
def __init__(
self,
email: str,
@@ -30,7 +28,7 @@ class LoginRequest:
class LoginResponse:
"""登录响应"""
def __init__(
self,
access_token: str,
@@ -52,18 +50,18 @@ class LoginResponse:
class LoginUseCase:
"""用户登录用例"""
def __init__(self, user_repository, session_store=None):
self.user_repository = user_repository
self.session_store = session_store or get_session_store()
def execute(self, request: LoginRequest) -> tuple[Optional[LoginResponse], Optional[str]]:
"""
执行登录
Args:
request: 登录请求
Returns:
(登录响应, 错误信息)
"""
@@ -71,31 +69,32 @@ class LoginUseCase:
# 1. 验证输入
if not request.email:
return None, "Email is required"
if not request.password:
return None, "Password is required"
# 2. 查找用户
user = self.user_repository.find_by_email(request.email)
if not user:
return None, "Invalid email or password"
# 3. 验证密码
if not password_hasher.verify_password(request.password, user.password_hash):
return None, "Invalid email or password"
# 4. 检查邮箱是否已验证(可选,根据需求决定是否强制)
# if not user.email_verified:
# return None, "Please verify your email first"
# 5. 创建 session 并生成 refresh_token
session_id = secrets.token_urlsafe(16)
refresh_token = secrets.token_urlsafe(32)
# 6. 生成基础 JWT token(包含 session_id,不包含 workspace
# 这里使用一个特殊的 "user_token",不包含 workspace 和 role
# 用户选择工作空间后,会换取包含 workspace 的 access_token
import jwt as pyjwt
now = datetime.now(timezone.utc)
access_token_payload = {
"sub": user.id,
@@ -107,7 +106,7 @@ class LoginUseCase:
access_token = pyjwt.encode(
access_token_payload,
jwt_service.config.SECRET_KEY,
algorithm=jwt_service.config.ALGORITHM
algorithm=jwt_service.config.ALGORITHM,
)
self.session_store.save_session(
session_id=session_id,
@@ -117,74 +116,77 @@ class LoginUseCase:
ip_address=request.ip_address,
expires_in_seconds=30 * 24 * 3600, # 30 天
)
# 8. 更新最后登录信息
user.last_login_at = datetime.now(timezone.utc)
user.last_login_ip = request.ip_address
self.user_repository.save(user)
# 9. 返回响应
return LoginResponse(
access_token=access_token,
refresh_token=refresh_token,
user_id=user.id,
email=user.email,
username=user.username,
display_name=user.display_name,
expires_in=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
), None
return (
LoginResponse(
access_token=access_token,
refresh_token=refresh_token,
user_id=user.id,
email=user.email,
username=user.username,
display_name=user.display_name,
expires_in=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
),
None,
)
except Exception as e:
return None, f"Login failed: {str(e)}"
class RefreshTokenRequest:
"""刷新令牌请求"""
def __init__(self, refresh_token: str):
self.refresh_token = refresh_token
class RefreshTokenUseCase:
"""刷新令牌用例"""
def __init__(self, user_repository):
self.user_repository = user_repository
def execute(self, request: RefreshTokenRequest) -> tuple[Optional[LoginResponse], Optional[str]]:
"""
执行令牌刷新
Args:
request: 刷新请求
Returns:
(登录响应, 错误信息)
"""
try:
if not request.refresh_token:
return None, "Refresh token is required"
# 1. 查找 session(通过遍历所有 session
# 注意:这里为了简化,先用遍历实现,生产环境应该用 refresh_token -> session_id 的索引
session = None
session_id = None
# 这是一个简化实现,实际应该在 SessionStore 中添加 find_by_refresh_token 方法
# 这里我们假设 refresh_token 就是 session_id(简化处理)
# 生产环境需要更复杂的映射
# 临时方案:从 Redis 获取(需要在 session_store 中添加方法)
# 现在先返回错误,提示需要实现
return None, "Refresh token implementation pending (需要完善 session_store)"
except Exception as e:
return None, f"Token refresh failed: {str(e)}"
class LogoutRequest:
"""登出请求"""
def __init__(
self,
user_id: str,
@@ -198,17 +200,17 @@ class LogoutRequest:
class LogoutUseCase:
"""用户登出用例"""
def __init__(self, session_store=None):
self.session_store = session_store or get_session_store()
def execute(self, request: LogoutRequest) -> tuple[bool, Optional[str]]:
"""
执行登出
Args:
request: 登出请求
Returns:
(是否成功, 错误信息)
"""
@@ -221,12 +223,12 @@ class LogoutUseCase:
# 删除当前 session
if not request.session_id:
return False, "Session ID is required"
success = self.session_store.delete_session(request.session_id)
if success:
return True, None
else:
return False, "Session not found"
except Exception as e:
return False, f"Logout failed: {str(e)}"
@@ -1,6 +1,7 @@
"""
密码重置 Use Case
"""
import secrets
from datetime import datetime, timedelta, timezone
from typing import Optional
@@ -11,14 +12,14 @@ from packages.domain.auth import password_hasher, password_validator
class RequestPasswordResetRequest:
"""请求密码重置"""
def __init__(self, email: str):
self.email = email.strip().lower()
class RequestPasswordResetUseCase:
"""请求密码重置用例"""
def __init__(
self,
user_repository,
@@ -30,41 +31,39 @@ class RequestPasswordResetUseCase:
self.base_url = base_url
self.token_expire_hours = token_expire_hours
self.email_service = email_service or get_email_service()
def execute(self, request: RequestPasswordResetRequest) -> tuple[bool, Optional[str]]:
"""
执行密码重置请求
Args:
request: 重置请求
Returns:
(是否成功, 错误信息)
"""
try:
if not request.email:
return False, "Email is required"
# 查找用户
user = self.user_repository.find_by_email(request.email)
# 安全考虑:即使用户不存在,也返回成功(避免暴露用户存在性)
if not user:
return True, None
# 生成重置令牌
reset_token = secrets.token_urlsafe(32)
reset_url = f"{self.base_url}/reset-password?token={reset_token}"
# 设置令牌和过期时间
user.password_reset_token = reset_token
user.password_reset_expires_at = datetime.now(timezone.utc) + timedelta(
hours=self.token_expire_hours
)
user.password_reset_expires_at = datetime.now(timezone.utc) + timedelta(hours=self.token_expire_hours)
# 保存用户
self.user_repository.save(user)
# 发送重置邮件
try:
success, error = self.email_service.send_password_reset_email(
@@ -72,22 +71,22 @@ class RequestPasswordResetUseCase:
username=user.username or user.display_name,
reset_url=reset_url,
)
if not success:
print(f"Failed to send password reset email: {error}")
# 不返回错误,避免暴露用户存在性
except Exception as e:
print(f"Email service error: {e}")
return True, None
except Exception as e:
return False, f"Password reset request failed: {str(e)}"
class ResetPasswordRequest:
"""重置密码请求"""
def __init__(self, token: str, new_password: str):
self.token = token
self.new_password = new_password
@@ -95,54 +94,54 @@ class ResetPasswordRequest:
class ResetPasswordUseCase:
"""重置密码用例"""
def __init__(self, user_repository):
self.user_repository = user_repository
def execute(self, request: ResetPasswordRequest) -> tuple[bool, Optional[str]]:
"""
执行密码重置
Args:
request: 重置请求
Returns:
(是否成功, 错误信息)
"""
try:
if not request.token:
return False, "Reset token is required"
if not request.new_password:
return False, "New password is required"
# 验证新密码强度
valid, error = password_validator.validate(request.new_password)
if not valid:
return False, error
# 查找用户
user = self.user_repository.find_by_password_reset_token(request.token)
if not user:
return False, "Invalid or expired reset token"
# 检查令牌是否过期
if user.password_reset_expires_at:
if datetime.now(timezone.utc) > user.password_reset_expires_at:
return False, "Reset token has expired"
# 哈希新密码
hashed_password = password_hasher.hash_password(request.new_password)
# 更新用户密码
user.password_hash = hashed_password
user.password_reset_token = None
user.password_reset_expires_at = None
# 保存用户
self.user_repository.save(user)
return True, None
except Exception as e:
return False, f"Password reset failed: {str(e)}"
@@ -1,19 +1,20 @@
"""
用户注册 Use Case
"""
import secrets
from datetime import datetime, timedelta, timezone
from typing import Optional
from uuid import uuid4
from packages.adapters.smtp import get_email_service
from packages.domain.entities import User
from packages.domain.auth import password_hasher, password_validator
from packages.domain.entities import User
class RegisterUserRequest:
"""注册请求"""
def __init__(
self,
email: str,
@@ -29,7 +30,7 @@ class RegisterUserRequest:
class RegisterUserResponse:
"""注册响应"""
def __init__(
self,
user_id: str,
@@ -47,7 +48,7 @@ class RegisterUserResponse:
class RegisterUserUseCase:
"""用户注册用例"""
def __init__(
self,
user_repository,
@@ -56,7 +57,7 @@ class RegisterUserUseCase:
):
"""
初始化注册用例
Args:
user_repository: 用户仓储
base_url: 应用基础 URL(用于生成验证链接)
@@ -64,14 +65,14 @@ class RegisterUserUseCase:
self.user_repository = user_repository
self.base_url = base_url
self.email_service = email_service or get_email_service()
def execute(self, request: RegisterUserRequest) -> tuple[Optional[RegisterUserResponse], Optional[str]]:
"""
执行注册
Args:
request: 注册请求
Returns:
(注册响应, 错误信息)
"""
@@ -79,34 +80,34 @@ class RegisterUserUseCase:
# 1. 验证输入
if not request.email:
return None, "Email is required"
if not request.username:
return None, "Username is required"
if not request.display_name:
return None, "Display name is required"
# 2. 验证密码强度
valid, error = password_validator.validate(request.password)
if not valid:
return None, error
# 3. 检查邮箱是否已存在
existing_user = self.user_repository.find_by_email(request.email)
if existing_user:
return None, "Email already registered"
# 4. 检查用户名是否已存在
existing_username = self.user_repository.find_by_username(request.username)
if existing_username:
return None, "Username already taken"
# 5. 哈希密码
hashed_password = password_hasher.hash_password(request.password)
# 6. 生成邮箱验证令牌
verification_token = secrets.token_urlsafe(32)
# 7. 创建用户
user = User(
id=uuid4().hex,
@@ -118,14 +119,14 @@ class RegisterUserUseCase:
email_verification_token=verification_token,
created_at=datetime.now(timezone.utc),
)
# 8. 保存用户
self.user_repository.save(user)
# 9. 发送验证邮件
verification_url = f"{self.base_url}/verify-email?token={verification_token}"
email_sent = False
try:
success, error = self.email_service.send_verification_email(
to_email=user.email,
@@ -133,68 +134,71 @@ class RegisterUserUseCase:
verification_url=verification_url,
)
email_sent = success
if not success:
print(f"Failed to send verification email: {error}")
except Exception as e:
print(f"Email service error: {e}")
# 10. 返回响应(即使邮件发送失败,用户也已创建)
return RegisterUserResponse(
user_id=user.id,
email=user.email,
username=user.username,
display_name=user.display_name,
email_verification_sent=email_sent,
), None
return (
RegisterUserResponse(
user_id=user.id,
email=user.email,
username=user.username,
display_name=user.display_name,
email_verification_sent=email_sent,
),
None,
)
except Exception as e:
return None, f"Registration failed: {str(e)}"
class VerifyEmailRequest:
"""邮箱验证请求"""
def __init__(self, token: str):
self.token = token
class VerifyEmailUseCase:
"""邮箱验证用例"""
def __init__(self, user_repository):
self.user_repository = user_repository
def execute(self, request: VerifyEmailRequest) -> tuple[bool, Optional[str]]:
"""
执行邮箱验证
Args:
request: 验证请求
Returns:
(是否成功, 错误信息)
"""
try:
if not request.token:
return False, "Verification token is required"
# 查找用户
user = self.user_repository.find_by_verification_token(request.token)
if not user:
return False, "Invalid or expired verification token"
# 检查是否已验证
if user.email_verified:
return True, None # 已验证,返回成功
# 更新用户状态
user.email_verified = True
user.email_verification_token = None # 清空令牌
self.user_repository.save(user)
return True, None
except Exception as e:
return False, f"Email verification failed: {str(e)}"
+15 -11
View File
@@ -1,24 +1,26 @@
"""
通用分页器
"""
from typing import Generic, TypeVar, List, Optional
from pydantic import BaseModel, Field
from math import ceil
from math import ceil
from typing import Generic, List, Optional, TypeVar
from pydantic import BaseModel, Field
T = TypeVar("T")
class PaginationParams(BaseModel):
"""分页参数"""
page: int = Field(1, ge=1, description="页码(从 1 开始)")
page_size: int = Field(20, ge=1, le=100, description="每页数量(最大 100")
@property
def offset(self) -> int:
"""计算偏移量"""
return (self.page - 1) * self.page_size
@property
def limit(self) -> int:
"""返回限制数量"""
@@ -27,13 +29,14 @@ class PaginationParams(BaseModel):
class PaginationMeta(BaseModel):
"""分页元数据"""
page: int = Field(..., description="当前页码")
page_size: int = Field(..., description="每页数量")
total: int = Field(..., description="总记录数")
total_pages: int = Field(..., description="总页数")
has_next: bool = Field(..., description="是否有下一页")
has_prev: bool = Field(..., description="是否有上一页")
@classmethod
def from_params(
cls,
@@ -42,7 +45,7 @@ class PaginationMeta(BaseModel):
) -> "PaginationMeta":
"""从参数和总数创建元数据"""
total_pages = ceil(total / params.page_size) if total > 0 else 0
return cls(
page=params.page,
page_size=params.page_size,
@@ -55,9 +58,10 @@ class PaginationMeta(BaseModel):
class PaginatedResponse(BaseModel, Generic[T]):
"""分页响应"""
data: List[T] = Field(..., description="数据列表")
pagination: PaginationMeta = Field(..., description="分页信息")
@classmethod
def create(
cls,
@@ -78,11 +82,11 @@ def paginate(
) -> PaginatedResponse[T]:
"""
内存分页(适用于 InMemory Repository
Args:
items: 完整列表
params: 分页参数
Returns:
分页响应
"""
@@ -90,7 +94,7 @@ def paginate(
start = params.offset
end = start + params.limit
page_data = items[start:end]
return PaginatedResponse.create(
data=page_data,
params=params,
@@ -1,4 +1,5 @@
"""获取单个任务详情用例"""
from packages.domain import Task
from packages.ports import TaskRepository
@@ -1,4 +1,5 @@
"""项目管理 Use Cases"""
from packages.domain import Milestone, Task, TaskIssue, TaskPriority, TaskStatus
from packages.ports import MilestoneRepository, TaskIssueRepository, TaskRepository
@@ -1,4 +1,5 @@
"""更新任务基本信息用例"""
from packages.domain import Task
from packages.ports import TaskRepository
+35 -34
View File
@@ -1,53 +1,54 @@
"""Workspace 相关 Use Cases"""
from packages.application.workspace.create_workspace_use_case import (
CreateWorkspaceUseCase,
CreateWorkspaceRequest,
CreateWorkspaceResponse,
)
from packages.application.workspace.invite_member_use_case import (
InviteMemberUseCase,
InviteMemberRequest,
InviteMemberResponse,
)
from packages.application.workspace.accept_invitation_use_case import (
AcceptInvitationUseCase,
AcceptInvitationRequest,
AcceptInvitationResponse,
DeclineInvitationUseCase,
AcceptInvitationUseCase,
DeclineInvitationRequest,
DeclineInvitationUseCase,
)
from packages.application.workspace.remove_member_use_case import (
RemoveMemberUseCase,
RemoveMemberRequest,
LeaveWorkspaceUseCase,
LeaveWorkspaceRequest,
from packages.application.workspace.create_workspace_use_case import (
CreateWorkspaceRequest,
CreateWorkspaceResponse,
CreateWorkspaceUseCase,
)
from packages.application.workspace.update_member_role_use_case import (
UpdateMemberRoleUseCase,
UpdateMemberRoleRequest,
UpdateMemberRoleResponse,
)
from packages.application.workspace.list_workspaces_use_case import (
ListWorkspacesUseCase,
ListWorkspacesRequest,
ListWorkspacesResponse,
GetWorkspaceDetailUseCase,
GetWorkspaceDetailRequest,
WorkspaceInfo,
WorkspaceDetailInfo,
from packages.application.workspace.invite_member_use_case import (
InviteMemberRequest,
InviteMemberResponse,
InviteMemberUseCase,
)
from packages.application.workspace.list_members_use_case import (
ListMembersUseCase,
ListMembersRequest,
ListMembersResponse,
ListMembersUseCase,
MemberInfo,
)
from packages.application.workspace.list_workspaces_use_case import (
GetWorkspaceDetailRequest,
GetWorkspaceDetailUseCase,
ListWorkspacesRequest,
ListWorkspacesResponse,
ListWorkspacesUseCase,
WorkspaceDetailInfo,
WorkspaceInfo,
)
from packages.application.workspace.remove_member_use_case import (
LeaveWorkspaceRequest,
LeaveWorkspaceUseCase,
RemoveMemberRequest,
RemoveMemberUseCase,
)
from packages.application.workspace.subscription_use_case import (
UpgradeSubscriptionUseCase,
CancelSubscriptionRequest,
CancelSubscriptionUseCase,
UpgradeSubscriptionRequest,
UpgradeSubscriptionResponse,
CancelSubscriptionUseCase,
CancelSubscriptionRequest,
UpgradeSubscriptionUseCase,
)
from packages.application.workspace.update_member_role_use_case import (
UpdateMemberRoleRequest,
UpdateMemberRoleResponse,
UpdateMemberRoleUseCase,
)
__all__ = [
@@ -1,19 +1,17 @@
"""
接受/拒绝邀请 Use Case
"""
from datetime import datetime, timezone
from typing import Optional
from uuid import uuid4
from packages.domain.entities import (
WorkspaceMember,
InvitationStatus,
)
from packages.domain.entities import InvitationStatus, WorkspaceMember
class AcceptInvitationRequest:
"""接受邀请请求"""
def __init__(self, invitation_token: str, user_id: str):
self.invitation_token = invitation_token
self.user_id = user_id
@@ -21,7 +19,7 @@ class AcceptInvitationRequest:
class AcceptInvitationResponse:
"""接受邀请响应"""
def __init__(
self,
workspace_id: str,
@@ -35,7 +33,7 @@ class AcceptInvitationResponse:
class AcceptInvitationUseCase:
"""接受邀请用例"""
def __init__(
self,
workspace_repository,
@@ -47,14 +45,14 @@ class AcceptInvitationUseCase:
self.workspace_member_repository = workspace_member_repository
self.workspace_invitation_repository = workspace_invitation_repository
self.user_repository = user_repository
def execute(self, request: AcceptInvitationRequest) -> tuple[Optional[AcceptInvitationResponse], Optional[str]]:
"""
执行接受邀请
Args:
request: 接受请求
Returns:
(响应, 错误信息)
"""
@@ -62,40 +60,40 @@ class AcceptInvitationUseCase:
# 1. 验证输入
if not request.invitation_token:
return None, "Invitation token is required"
if not request.user_id:
return None, "User ID is required"
# 2. 查找邀请
invitation = self.workspace_invitation_repository.find_by_token(request.invitation_token)
if not invitation:
return None, "Invalid invitation token"
# 3. 检查邀请状态
if invitation.status != InvitationStatus.PENDING:
return None, f"Invitation has already been {invitation.status}"
# 4. 检查是否过期
if invitation.expires_at and datetime.now(timezone.utc) > invitation.expires_at:
# 更新状态为过期
invitation.status = InvitationStatus.EXPIRED
self.workspace_invitation_repository.save(invitation)
return None, "Invitation has expired"
# 5. 验证用户存在
user = self.user_repository.find_by_id(request.user_id)
if not user:
return None, "User not found"
# 6. 验证用户邮箱匹配
if user.email.lower() != invitation.invitee_email.lower():
return None, "This invitation is for a different email address"
# 7. 验证 Workspace 存在
workspace = self.workspace_repository.find_by_id(invitation.workspace_id)
if not workspace:
return None, "Workspace not found"
# 8. 检查用户是否已经是成员
existing_member = self.workspace_member_repository.find_by_workspace_and_user(
invitation.workspace_id,
@@ -106,13 +104,16 @@ class AcceptInvitationUseCase:
invitation.status = InvitationStatus.ACCEPTED
invitation.accepted_at = datetime.now(timezone.utc)
self.workspace_invitation_repository.save(invitation)
return AcceptInvitationResponse(
workspace_id=workspace.id,
workspace_name=workspace.name,
role=existing_member.role,
), None
return (
AcceptInvitationResponse(
workspace_id=workspace.id,
workspace_name=workspace.name,
role=existing_member.role,
),
None,
)
# 9. 创建成员记录
member = WorkspaceMember(
id=uuid4().hex,
@@ -122,45 +123,48 @@ class AcceptInvitationUseCase:
invited_by=invitation.inviter_user_id,
joined_at=datetime.now(timezone.utc),
)
self.workspace_member_repository.save(member)
# 10. 更新邀请状态
invitation.status = InvitationStatus.ACCEPTED
invitation.accepted_at = datetime.now(timezone.utc)
self.workspace_invitation_repository.save(invitation)
# 11. 返回响应
return AcceptInvitationResponse(
workspace_id=workspace.id,
workspace_name=workspace.name,
role=member.role,
), None
return (
AcceptInvitationResponse(
workspace_id=workspace.id,
workspace_name=workspace.name,
role=member.role,
),
None,
)
except Exception as e:
return None, f"Failed to accept invitation: {str(e)}"
class DeclineInvitationRequest:
"""拒绝邀请请求"""
def __init__(self, invitation_token: str):
self.invitation_token = invitation_token
class DeclineInvitationUseCase:
"""拒绝邀请用例"""
def __init__(self, workspace_invitation_repository):
self.workspace_invitation_repository = workspace_invitation_repository
def execute(self, request: DeclineInvitationRequest) -> tuple[bool, Optional[str]]:
"""
执行拒绝邀请
Args:
request: 拒绝请求
Returns:
(是否成功, 错误信息)
"""
@@ -168,21 +172,21 @@ class DeclineInvitationUseCase:
# 1. 验证输入
if not request.invitation_token:
return False, "Invitation token is required"
# 2. 查找邀请
invitation = self.workspace_invitation_repository.find_by_token(request.invitation_token)
if not invitation:
return False, "Invalid invitation token"
# 3. 检查邀请状态
if invitation.status != InvitationStatus.PENDING:
return False, f"Invitation has already been {invitation.status}"
# 4. 更新状态为已拒绝
invitation.status = InvitationStatus.DECLINED
self.workspace_invitation_repository.save(invitation)
return True, None
except Exception as e:
return False, f"Failed to decline invitation: {str(e)}"
@@ -1,16 +1,17 @@
"""
创建 Workspace Use Case
"""
from datetime import datetime, timezone
from typing import Optional
from uuid import uuid4
from datetime import datetime, timezone
from packages.domain.entities import Workspace, WorkspaceMember, WorkspaceMemberRole
class CreateWorkspaceRequest:
"""创建工作空间请求"""
def __init__(
self,
name: str,
@@ -24,7 +25,7 @@ class CreateWorkspaceRequest:
class CreateWorkspaceResponse:
"""创建工作空间响应"""
def __init__(
self,
workspace_id: str,
@@ -42,14 +43,14 @@ class CreateWorkspaceResponse:
class CreateWorkspaceUseCase:
"""创建工作空间用例"""
# 订阅计划配额配置
PLAN_QUOTAS = {
"free": {"max_projects": 3, "max_storage_gb": 10},
"pro": {"max_projects": 999999, "max_storage_gb": 100}, # 999999 表示无限
"enterprise": {"max_projects": 999999, "max_storage_gb": 1000},
}
def __init__(
self,
workspace_repository,
@@ -59,14 +60,14 @@ class CreateWorkspaceUseCase:
self.workspace_repository = workspace_repository
self.workspace_member_repository = workspace_member_repository
self.user_repository = user_repository
def execute(self, request: CreateWorkspaceRequest) -> tuple[Optional[CreateWorkspaceResponse], Optional[str]]:
"""
执行创建工作空间
Args:
request: 创建请求
Returns:
(响应, 错误信息)
"""
@@ -74,25 +75,25 @@ class CreateWorkspaceUseCase:
# 1. 验证输入
if not request.name:
return None, "Workspace name is required"
if len(request.name) > 100:
return None, "Workspace name is too long (max 100 characters)"
if not request.owner_user_id:
return None, "Owner user ID is required"
# 2. 验证用户存在
owner = self.user_repository.find_by_id(request.owner_user_id)
if not owner:
return None, "Owner user not found"
# 3. 验证订阅计划
if request.subscription_plan not in self.PLAN_QUOTAS:
return None, f"Invalid subscription plan: {request.subscription_plan}"
# 4. 获取配额
quota = self.PLAN_QUOTAS[request.subscription_plan]
# 5. 创建 Workspace
workspace = Workspace(
id=uuid4().hex,
@@ -105,10 +106,10 @@ class CreateWorkspaceUseCase:
used_storage_gb=0.0,
created_at=datetime.now(timezone.utc),
)
# 6. 保存 Workspace
self.workspace_repository.save(workspace)
# 7. 创建 Owner 成员记录
owner_member = WorkspaceMember(
id=uuid4().hex,
@@ -118,17 +119,20 @@ class CreateWorkspaceUseCase:
invited_by=None, # Owner 不需要邀请
joined_at=datetime.now(timezone.utc),
)
self.workspace_member_repository.save(owner_member)
# 8. 返回响应
return CreateWorkspaceResponse(
workspace_id=workspace.id,
name=workspace.name,
subscription_plan=workspace.subscription_plan,
max_projects=workspace.max_projects,
max_storage_gb=workspace.max_storage_gb,
), None
return (
CreateWorkspaceResponse(
workspace_id=workspace.id,
name=workspace.name,
subscription_plan=workspace.subscription_plan,
max_projects=workspace.max_projects,
max_storage_gb=workspace.max_storage_gb,
),
None,
)
except Exception as e:
return None, f"Failed to create workspace: {str(e)}"
@@ -1,6 +1,7 @@
"""
邀请成员到 Workspace Use Case
"""
import secrets
from datetime import datetime, timedelta, timezone
from typing import Optional
@@ -8,15 +9,15 @@ from uuid import uuid4
from packages.adapters.smtp import get_email_service
from packages.domain.entities import (
InvitationStatus,
WorkspaceInvitation,
WorkspaceMemberRole,
InvitationStatus,
)
class InviteMemberRequest:
"""邀请成员请求"""
def __init__(
self,
workspace_id: str,
@@ -32,7 +33,7 @@ class InviteMemberRequest:
class InviteMemberResponse:
"""邀请成员响应"""
def __init__(
self,
invitation_id: str,
@@ -48,13 +49,13 @@ class InviteMemberResponse:
class InviteMemberUseCase:
"""邀请成员用例"""
VALID_ROLES = [
WorkspaceMemberRole.ADMIN,
WorkspaceMemberRole.MEMBER,
WorkspaceMemberRole.VIEWER,
]
def __init__(
self,
workspace_repository,
@@ -72,14 +73,14 @@ class InviteMemberUseCase:
self.base_url = base_url
self.invitation_expire_days = invitation_expire_days
self.email_service = email_service or get_email_service()
def execute(self, request: InviteMemberRequest) -> tuple[Optional[InviteMemberResponse], Optional[str]]:
"""
执行邀请成员
Args:
request: 邀请请求
Returns:
(响应, 错误信息)
"""
@@ -87,25 +88,25 @@ class InviteMemberUseCase:
# 1. 验证输入
if not request.workspace_id:
return None, "Workspace ID is required"
if not request.inviter_user_id:
return None, "Inviter user ID is required"
if not request.invitee_email:
return None, "Invitee email is required"
if not request.role:
return None, "Role is required"
# 2. 验证角色(不能邀请 owner)
if request.role not in self.VALID_ROLES:
return None, f"Invalid role: {request.role}. Cannot invite as owner."
# 3. 验证 Workspace 存在
workspace = self.workspace_repository.find_by_id(request.workspace_id)
if not workspace:
return None, "Workspace not found"
# 4. 验证邀请人是成员且有权限(owner 或 admin
inviter_member = self.workspace_member_repository.find_by_workspace_and_user(
request.workspace_id,
@@ -113,10 +114,13 @@ class InviteMemberUseCase:
)
if not inviter_member:
return None, "You are not a member of this workspace"
if inviter_member.role not in [WorkspaceMemberRole.OWNER, WorkspaceMemberRole.ADMIN]:
if inviter_member.role not in [
WorkspaceMemberRole.OWNER,
WorkspaceMemberRole.ADMIN,
]:
return None, "Only owners and admins can invite members"
# 5. 检查被邀请人是否已经是成员
invitee_user = self.user_repository.find_by_email(request.invitee_email)
if invitee_user:
@@ -126,7 +130,7 @@ class InviteMemberUseCase:
)
if existing_member:
return None, "User is already a member of this workspace"
# 6. 检查是否已有待处理的邀请
existing_invitation = self.workspace_invitation_repository.find_pending_by_workspace_and_email(
request.workspace_id,
@@ -134,11 +138,11 @@ class InviteMemberUseCase:
)
if existing_invitation:
return None, "An invitation has already been sent to this email"
# 7. 生成邀请令牌
invitation_token = secrets.token_urlsafe(32)
expires_at = datetime.now(timezone.utc) + timedelta(days=self.invitation_expire_days)
# 8. 创建邀请记录
invitation = WorkspaceInvitation(
id=uuid4().hex,
@@ -151,17 +155,17 @@ class InviteMemberUseCase:
expires_at=expires_at,
created_at=datetime.now(timezone.utc),
)
# 9. 保存邀请
self.workspace_invitation_repository.save(invitation)
# 10. 发送邀请邮件
invitation_url = f"{self.base_url}/invitations/{invitation_token}/accept"
try:
inviter = self.user_repository.find_by_id(request.inviter_user_id)
inviter_name = inviter.display_name if inviter else "Someone"
success, error = self.email_service.send_workspace_invitation_email(
to_email=request.invitee_email,
inviter_name=inviter_name,
@@ -169,19 +173,22 @@ class InviteMemberUseCase:
role=request.role,
invitation_url=invitation_url,
)
if not success:
print(f"Failed to send invitation email: {error}")
except Exception as e:
print(f"Email service error: {e}")
# 11. 返回响应
return InviteMemberResponse(
invitation_id=invitation.id,
invitee_email=invitation.invitee_email,
role=invitation.role,
expires_at=invitation.expires_at,
), None
return (
InviteMemberResponse(
invitation_id=invitation.id,
invitee_email=invitation.invitee_email,
role=invitation.role,
expires_at=invitation.expires_at,
),
None,
)
except Exception as e:
return None, f"Failed to invite member: {str(e)}"
@@ -1,13 +1,14 @@
"""
获取成员列表 Use Case
"""
from typing import Optional, List
from datetime import datetime
from typing import List, Optional
class MemberInfo:
"""成员信息"""
def __init__(
self,
member_id: str,
@@ -31,7 +32,7 @@ class MemberInfo:
class ListMembersRequest:
"""获取成员列表请求"""
def __init__(self, workspace_id: str, requester_user_id: str):
self.workspace_id = workspace_id
self.requester_user_id = requester_user_id
@@ -39,14 +40,14 @@ class ListMembersRequest:
class ListMembersResponse:
"""获取成员列表响应"""
def __init__(self, members: List[MemberInfo]):
self.members = members
class ListMembersUseCase:
"""获取成员列表用例"""
def __init__(
self,
workspace_repository,
@@ -56,14 +57,14 @@ class ListMembersUseCase:
self.workspace_repository = workspace_repository
self.workspace_member_repository = workspace_member_repository
self.user_repository = user_repository
def execute(self, request: ListMembersRequest) -> tuple[Optional[ListMembersResponse], Optional[str]]:
"""
执行获取成员列表
Args:
request: 请求
Returns:
(响应, 错误信息)
"""
@@ -71,15 +72,15 @@ class ListMembersUseCase:
# 1. 验证输入
if not request.workspace_id:
return None, "Workspace ID is required"
if not request.requester_user_id:
return None, "Requester user ID is required"
# 2. 验证工作空间存在
workspace = self.workspace_repository.find_by_id(request.workspace_id)
if not workspace:
return None, "Workspace not found"
# 3. 验证请求者是成员(只有成员才能查看成员列表)
requester_member = self.workspace_member_repository.find_by_workspace_and_user(
request.workspace_id,
@@ -87,17 +88,17 @@ class ListMembersUseCase:
)
if not requester_member:
return None, "You are not a member of this workspace"
# 4. 获取所有成员
members = self.workspace_member_repository.find_by_workspace(request.workspace_id)
# 5. 获取每个成员的用户信息
member_infos = []
for member in members:
user = self.user_repository.find_by_id(member.user_id)
if not user:
continue # 跳过不存在的用户
member_info = MemberInfo(
member_id=member.id,
user_id=user.id,
@@ -109,9 +110,9 @@ class ListMembersUseCase:
joined_at=member.joined_at,
)
member_infos.append(member_info)
# 6. 返回响应
return ListMembersResponse(members=member_infos), None
except Exception as e:
return None, f"Failed to list members: {str(e)}"
@@ -1,13 +1,14 @@
"""
获取工作空间列表和详情 Use Case
"""
from typing import Optional, List
from datetime import datetime
from typing import List, Optional
class WorkspaceInfo:
"""工作空间信息"""
def __init__(
self,
workspace_id: str,
@@ -37,21 +38,21 @@ class WorkspaceInfo:
class ListWorkspacesRequest:
"""获取工作空间列表请求"""
def __init__(self, user_id: str):
self.user_id = user_id
class ListWorkspacesResponse:
"""获取工作空间列表响应"""
def __init__(self, workspaces: List[WorkspaceInfo]):
self.workspaces = workspaces
class ListWorkspacesUseCase:
"""获取工作空间列表用例"""
def __init__(
self,
workspace_repository,
@@ -59,14 +60,14 @@ class ListWorkspacesUseCase:
):
self.workspace_repository = workspace_repository
self.workspace_member_repository = workspace_member_repository
def execute(self, request: ListWorkspacesRequest) -> tuple[Optional[ListWorkspacesResponse], Optional[str]]:
"""
执行获取工作空间列表
Args:
request: 请求
Returns:
(响应, 错误信息)
"""
@@ -74,20 +75,20 @@ class ListWorkspacesUseCase:
# 1. 验证输入
if not request.user_id:
return None, "User ID is required"
# 2. 获取用户所有的成员记录
memberships = self.workspace_member_repository.find_by_user(request.user_id)
# 3. 获取每个工作空间的信息
workspace_infos = []
for membership in memberships:
workspace = self.workspace_repository.find_by_id(membership.workspace_id)
if not workspace:
continue
# 获取成员数量
member_count = self.workspace_member_repository.count_by_workspace(membership.workspace_id)
workspace_info = WorkspaceInfo(
workspace_id=workspace.id,
name=workspace.name,
@@ -102,17 +103,17 @@ class ListWorkspacesUseCase:
created_at=workspace.created_at,
)
workspace_infos.append(workspace_info)
# 4. 返回响应
return ListWorkspacesResponse(workspaces=workspace_infos), None
except Exception as e:
return None, f"Failed to list workspaces: {str(e)}"
class GetWorkspaceDetailRequest:
"""获取工作空间详情请求"""
def __init__(self, workspace_id: str, user_id: str):
self.workspace_id = workspace_id
self.user_id = user_id
@@ -120,7 +121,7 @@ class GetWorkspaceDetailRequest:
class WorkspaceDetailInfo:
"""工作空间详情信息"""
def __init__(
self,
workspace_id: str,
@@ -152,7 +153,7 @@ class WorkspaceDetailInfo:
class GetWorkspaceDetailUseCase:
"""获取工作空间详情用例"""
def __init__(
self,
workspace_repository,
@@ -160,14 +161,14 @@ class GetWorkspaceDetailUseCase:
):
self.workspace_repository = workspace_repository
self.workspace_member_repository = workspace_member_repository
def execute(self, request: GetWorkspaceDetailRequest) -> tuple[Optional[WorkspaceDetailInfo], Optional[str]]:
"""
执行获取工作空间详情
Args:
request: 请求
Returns:
(详情信息, 错误信息)
"""
@@ -175,15 +176,15 @@ class GetWorkspaceDetailUseCase:
# 1. 验证输入
if not request.workspace_id:
return None, "Workspace ID is required"
if not request.user_id:
return None, "User ID is required"
# 2. 验证工作空间存在
workspace = self.workspace_repository.find_by_id(request.workspace_id)
if not workspace:
return None, "Workspace not found"
# 3. 验证用户是成员
membership = self.workspace_member_repository.find_by_workspace_and_user(
request.workspace_id,
@@ -191,10 +192,10 @@ class GetWorkspaceDetailUseCase:
)
if not membership:
return None, "You are not a member of this workspace"
# 4. 获取成员数量
member_count = self.workspace_member_repository.count_by_workspace(request.workspace_id)
# 5. 构建详情信息
detail_info = WorkspaceDetailInfo(
workspace_id=workspace.id,
@@ -210,8 +211,8 @@ class GetWorkspaceDetailUseCase:
user_role=membership.role,
created_at=workspace.created_at,
)
return detail_info, None
except Exception as e:
return None, f"Failed to get workspace detail: {str(e)}"
@@ -1,6 +1,7 @@
"""
移除成员 Use Case
"""
from typing import Optional
from packages.domain.entities import WorkspaceMemberRole
@@ -8,7 +9,7 @@ from packages.domain.entities import WorkspaceMemberRole
class RemoveMemberRequest:
"""移除成员请求"""
def __init__(
self,
workspace_id: str,
@@ -22,7 +23,7 @@ class RemoveMemberRequest:
class RemoveMemberUseCase:
"""移除成员用例"""
def __init__(
self,
workspace_repository,
@@ -30,14 +31,14 @@ class RemoveMemberUseCase:
):
self.workspace_repository = workspace_repository
self.workspace_member_repository = workspace_member_repository
def execute(self, request: RemoveMemberRequest) -> tuple[bool, Optional[str]]:
"""
执行移除成员
Args:
request: 移除请求
Returns:
(是否成功, 错误信息)
"""
@@ -45,18 +46,18 @@ class RemoveMemberUseCase:
# 1. 验证输入
if not request.workspace_id:
return False, "Workspace ID is required"
if not request.requester_user_id:
return False, "Requester user ID is required"
if not request.target_user_id:
return False, "Target user ID is required"
# 2. 验证 Workspace 存在
workspace = self.workspace_repository.find_by_id(request.workspace_id)
if not workspace:
return False, "Workspace not found"
# 3. 验证请求者是成员且有权限
requester_member = self.workspace_member_repository.find_by_workspace_and_user(
request.workspace_id,
@@ -64,10 +65,13 @@ class RemoveMemberUseCase:
)
if not requester_member:
return False, "You are not a member of this workspace"
if requester_member.role not in [WorkspaceMemberRole.OWNER, WorkspaceMemberRole.ADMIN]:
if requester_member.role not in [
WorkspaceMemberRole.OWNER,
WorkspaceMemberRole.ADMIN,
]:
return False, "Only owners and admins can remove members"
# 4. 验证目标成员存在
target_member = self.workspace_member_repository.find_by_workspace_and_user(
request.workspace_id,
@@ -75,34 +79,33 @@ class RemoveMemberUseCase:
)
if not target_member:
return False, "Target user is not a member of this workspace"
# 5. 不能移除自己(应该用离开 workspace 的功能)
if request.requester_user_id == request.target_user_id:
return False, "Cannot remove yourself. Use leave workspace instead."
# 6. 不能移除 owner
if target_member.role == WorkspaceMemberRole.OWNER:
return False, "Cannot remove the workspace owner"
# 7. Admin 不能移除另一个 Admin(只有 owner 可以)
if (requester_member.role == WorkspaceMemberRole.ADMIN and
target_member.role == WorkspaceMemberRole.ADMIN):
if requester_member.role == WorkspaceMemberRole.ADMIN and target_member.role == WorkspaceMemberRole.ADMIN:
return False, "Admins cannot remove other admins"
# 8. 删除成员记录
success = self.workspace_member_repository.delete(target_member.id)
if not success:
return False, "Failed to remove member"
return True, None
except Exception as e:
return False, f"Failed to remove member: {str(e)}"
class LeaveWorkspaceRequest:
"""离开 Workspace 请求"""
def __init__(self, workspace_id: str, user_id: str):
self.workspace_id = workspace_id
self.user_id = user_id
@@ -110,7 +113,7 @@ class LeaveWorkspaceRequest:
class LeaveWorkspaceUseCase:
"""离开 Workspace 用例"""
def __init__(
self,
workspace_repository,
@@ -118,14 +121,14 @@ class LeaveWorkspaceUseCase:
):
self.workspace_repository = workspace_repository
self.workspace_member_repository = workspace_member_repository
def execute(self, request: LeaveWorkspaceRequest) -> tuple[bool, Optional[str]]:
"""
执行离开 Workspace
Args:
request: 离开请求
Returns:
(是否成功, 错误信息)
"""
@@ -133,15 +136,15 @@ class LeaveWorkspaceUseCase:
# 1. 验证输入
if not request.workspace_id:
return False, "Workspace ID is required"
if not request.user_id:
return False, "User ID is required"
# 2. 验证 Workspace 存在
workspace = self.workspace_repository.find_by_id(request.workspace_id)
if not workspace:
return False, "Workspace not found"
# 3. 验证用户是成员
member = self.workspace_member_repository.find_by_workspace_and_user(
request.workspace_id,
@@ -149,17 +152,20 @@ class LeaveWorkspaceUseCase:
)
if not member:
return False, "You are not a member of this workspace"
# 4. Owner 不能离开(需要先转移 ownership 或删除 workspace
if member.role == WorkspaceMemberRole.OWNER:
return False, "Owner cannot leave workspace. Transfer ownership or delete workspace first."
return (
False,
"Owner cannot leave workspace. Transfer ownership or delete workspace first.",
)
# 5. 删除成员记录
success = self.workspace_member_repository.delete(member.id)
if not success:
return False, "Failed to leave workspace"
return True, None
except Exception as e:
return False, f"Failed to leave workspace: {str(e)}"
@@ -1,15 +1,16 @@
"""
Subscription 管理 Use Case
"""
from typing import Optional
from datetime import datetime, timedelta, timezone
from typing import Optional
from packages.domain.entities import WorkspaceMemberRole
class UpgradeSubscriptionRequest:
"""升级订阅请求"""
def __init__(
self,
workspace_id: str,
@@ -23,7 +24,7 @@ class UpgradeSubscriptionRequest:
class UpgradeSubscriptionResponse:
"""升级订阅响应"""
def __init__(
self,
workspace_id: str,
@@ -41,21 +42,21 @@ class UpgradeSubscriptionResponse:
class UpgradeSubscriptionUseCase:
"""升级订阅用例"""
# 订阅计划配额
PLAN_QUOTAS = {
"free": {"max_projects": 3, "max_storage_gb": 10, "price": 0},
"pro": {"max_projects": 999999, "max_storage_gb": 100, "price": 99},
"enterprise": {"max_projects": 999999, "max_storage_gb": 1000, "price": 999},
}
# 计划等级
PLAN_LEVELS = {
"free": 0,
"pro": 1,
"enterprise": 2,
}
def __init__(
self,
workspace_repository,
@@ -63,14 +64,16 @@ class UpgradeSubscriptionUseCase:
):
self.workspace_repository = workspace_repository
self.workspace_member_repository = workspace_member_repository
def execute(self, request: UpgradeSubscriptionRequest) -> tuple[Optional[UpgradeSubscriptionResponse], Optional[str]]:
def execute(
self, request: UpgradeSubscriptionRequest
) -> tuple[Optional[UpgradeSubscriptionResponse], Optional[str]]:
"""
执行升级订阅
Args:
request: 升级请求
Returns:
(响应, 错误信息)
"""
@@ -78,22 +81,22 @@ class UpgradeSubscriptionUseCase:
# 1. 验证输入
if not request.workspace_id:
return None, "Workspace ID is required"
if not request.requester_user_id:
return None, "Requester user ID is required"
if not request.new_plan:
return None, "New plan is required"
# 2. 验证新计划有效
if request.new_plan not in self.PLAN_QUOTAS:
return None, f"Invalid plan: {request.new_plan}"
# 3. 验证工作空间存在
workspace = self.workspace_repository.find_by_id(request.workspace_id)
if not workspace:
return None, "Workspace not found"
# 4. 验证权限(只有 Owner 可以管理订阅)
member = self.workspace_member_repository.find_by_workspace_and_user(
request.workspace_id,
@@ -101,50 +104,56 @@ class UpgradeSubscriptionUseCase:
)
if not member:
return None, "You are not a member of this workspace"
if member.role != WorkspaceMemberRole.OWNER:
return None, "Only workspace owner can manage subscription"
# 5. 检查是否是升级(不能降级到免费计划,需要用取消订阅)
current_level = self.PLAN_LEVELS.get(workspace.subscription_plan, 0)
new_level = self.PLAN_LEVELS.get(request.new_plan, 0)
if new_level < current_level:
return None, "Cannot downgrade plan. Use cancel subscription to return to free plan."
return (
None,
"Cannot downgrade plan. Use cancel subscription to return to free plan.",
)
if new_level == current_level:
return None, f"Workspace is already on {request.new_plan} plan"
# 6. 更新订阅
old_plan = workspace.subscription_plan
quota = self.PLAN_QUOTAS[request.new_plan]
workspace.subscription_plan = request.new_plan
workspace.subscription_status = "active"
workspace.max_projects = quota["max_projects"]
workspace.max_storage_gb = quota["max_storage_gb"]
# 设置过期时间(假设按月订阅)
workspace.subscription_expires_at = datetime.now(timezone.utc) + timedelta(days=30)
self.workspace_repository.save(workspace)
# 7. 返回响应
return UpgradeSubscriptionResponse(
workspace_id=workspace.id,
old_plan=old_plan,
new_plan=workspace.subscription_plan,
max_projects=workspace.max_projects,
max_storage_gb=workspace.max_storage_gb,
), None
return (
UpgradeSubscriptionResponse(
workspace_id=workspace.id,
old_plan=old_plan,
new_plan=workspace.subscription_plan,
max_projects=workspace.max_projects,
max_storage_gb=workspace.max_storage_gb,
),
None,
)
except Exception as e:
return None, f"Failed to upgrade subscription: {str(e)}"
class CancelSubscriptionRequest:
"""取消订阅请求"""
def __init__(self, workspace_id: str, requester_user_id: str):
self.workspace_id = workspace_id
self.requester_user_id = requester_user_id
@@ -152,7 +161,7 @@ class CancelSubscriptionRequest:
class CancelSubscriptionUseCase:
"""取消订阅用例"""
def __init__(
self,
workspace_repository,
@@ -160,14 +169,14 @@ class CancelSubscriptionUseCase:
):
self.workspace_repository = workspace_repository
self.workspace_member_repository = workspace_member_repository
def execute(self, request: CancelSubscriptionRequest) -> tuple[bool, Optional[str]]:
"""
执行取消订阅
Args:
request: 取消请求
Returns:
(是否成功, 错误信息)
"""
@@ -175,15 +184,15 @@ class CancelSubscriptionUseCase:
# 1. 验证输入
if not request.workspace_id:
return False, "Workspace ID is required"
if not request.requester_user_id:
return False, "Requester user ID is required"
# 2. 验证工作空间存在
workspace = self.workspace_repository.find_by_id(request.workspace_id)
if not workspace:
return False, "Workspace not found"
# 3. 验证权限(只有 Owner 可以管理订阅)
member = self.workspace_member_repository.find_by_workspace_and_user(
request.workspace_id,
@@ -191,24 +200,24 @@ class CancelSubscriptionUseCase:
)
if not member:
return False, "You are not a member of this workspace"
if member.role != WorkspaceMemberRole.OWNER:
return False, "Only workspace owner can manage subscription"
# 4. 检查当前计划
if workspace.subscription_plan == "free":
return False, "Workspace is already on free plan"
# 5. 降级到 free 计划
workspace.subscription_plan = "free"
workspace.subscription_status = "active"
workspace.subscription_expires_at = None
workspace.max_projects = 3
workspace.max_storage_gb = 10
self.workspace_repository.save(workspace)
return True, None
except Exception as e:
return False, f"Failed to cancel subscription: {str(e)}"
@@ -1,6 +1,7 @@
"""
修改成员角色 Use Case
"""
from typing import Optional
from packages.domain.entities import WorkspaceMemberRole
@@ -8,7 +9,7 @@ from packages.domain.entities import WorkspaceMemberRole
class UpdateMemberRoleRequest:
"""修改成员角色请求"""
def __init__(
self,
workspace_id: str,
@@ -24,7 +25,7 @@ class UpdateMemberRoleRequest:
class UpdateMemberRoleResponse:
"""修改成员角色响应"""
def __init__(self, user_id: str, old_role: str, new_role: str):
self.user_id = user_id
self.old_role = old_role
@@ -33,13 +34,13 @@ class UpdateMemberRoleResponse:
class UpdateMemberRoleUseCase:
"""修改成员角色用例"""
VALID_ROLES = [
WorkspaceMemberRole.ADMIN,
WorkspaceMemberRole.MEMBER,
WorkspaceMemberRole.VIEWER,
]
def __init__(
self,
workspace_repository,
@@ -47,14 +48,14 @@ class UpdateMemberRoleUseCase:
):
self.workspace_repository = workspace_repository
self.workspace_member_repository = workspace_member_repository
def execute(self, request: UpdateMemberRoleRequest) -> tuple[Optional[UpdateMemberRoleResponse], Optional[str]]:
"""
执行修改成员角色
Args:
request: 修改请求
Returns:
(响应, 错误信息)
"""
@@ -62,25 +63,28 @@ class UpdateMemberRoleUseCase:
# 1. 验证输入
if not request.workspace_id:
return None, "Workspace ID is required"
if not request.requester_user_id:
return None, "Requester user ID is required"
if not request.target_user_id:
return None, "Target user ID is required"
if not request.new_role:
return None, "New role is required"
# 2. 验证新角色(不能修改为 owner)
if request.new_role not in self.VALID_ROLES:
return None, f"Invalid role: {request.new_role}. Cannot change to owner."
return (
None,
f"Invalid role: {request.new_role}. Cannot change to owner.",
)
# 3. 验证 Workspace 存在
workspace = self.workspace_repository.find_by_id(request.workspace_id)
if not workspace:
return None, "Workspace not found"
# 4. 验证请求者是成员且有权限(只有 owner 和 admin 可以修改角色)
requester_member = self.workspace_member_repository.find_by_workspace_and_user(
request.workspace_id,
@@ -88,10 +92,13 @@ class UpdateMemberRoleUseCase:
)
if not requester_member:
return None, "You are not a member of this workspace"
if requester_member.role not in [WorkspaceMemberRole.OWNER, WorkspaceMemberRole.ADMIN]:
if requester_member.role not in [
WorkspaceMemberRole.OWNER,
WorkspaceMemberRole.ADMIN,
]:
return None, "Only owners and admins can change member roles"
# 5. 验证目标成员存在
target_member = self.workspace_member_repository.find_by_workspace_and_user(
request.workspace_id,
@@ -99,35 +106,37 @@ class UpdateMemberRoleUseCase:
)
if not target_member:
return None, "Target user is not a member of this workspace"
# 6. 不能修改自己的角色
if request.requester_user_id == request.target_user_id:
return None, "Cannot change your own role"
# 7. 不能修改 owner 的角色
if target_member.role == WorkspaceMemberRole.OWNER:
return None, "Cannot change the owner's role"
# 8. Admin 不能修改另一个 Admin 的角色(只有 owner 可以)
if (requester_member.role == WorkspaceMemberRole.ADMIN and
target_member.role == WorkspaceMemberRole.ADMIN):
if requester_member.role == WorkspaceMemberRole.ADMIN and target_member.role == WorkspaceMemberRole.ADMIN:
return None, "Admins cannot change other admins' roles"
# 9. 检查角色是否相同
if target_member.role == request.new_role:
return None, f"User already has the {request.new_role} role"
# 10. 更新角色
old_role = target_member.role
target_member.role = request.new_role
self.workspace_member_repository.save(target_member)
# 11. 返回响应
return UpdateMemberRoleResponse(
user_id=request.target_user_id,
old_role=old_role,
new_role=request.new_role,
), None
return (
UpdateMemberRoleResponse(
user_id=request.target_user_id,
old_role=old_role,
new_role=request.new_role,
),
None,
)
except Exception as e:
return None, f"Failed to update member role: {str(e)}"
+6 -2
View File
@@ -1,6 +1,10 @@
"""Domain package for core business entities and rules."""
from .classification import AssetClassification, ClassificationJob, ClassificationJobStatus
from .classification import (
AssetClassification,
ClassificationJob,
ClassificationJobStatus,
)
from .entities import (
Asset,
AssetLibrary,
@@ -13,8 +17,8 @@ from .entities import (
User,
Workspace,
)
from .generation_task import GenerationTask, GenerationTaskStatus
from .generated_video import GeneratedVideo
from .generation_task import GenerationTask, GenerationTaskStatus
from .project_management import Milestone, Task, TaskIssue, TaskPriority, TaskStatus
__all__ = [
+7 -2
View File
@@ -5,7 +5,13 @@ services such as Redis session storage and SMTP email delivery live under
`packages.adapters` and should be injected into use cases.
"""
from packages.domain.auth.jwt_service import JWTConfig, JWTService, TokenType, jwt_service
from packages.domain.auth.email_service import EmailConfig, EmailService
from packages.domain.auth.jwt_service import (
JWTConfig,
JWTService,
TokenType,
jwt_service,
)
from packages.domain.auth.password_hasher import (
PasswordHasher,
PasswordValidator,
@@ -13,7 +19,6 @@ from packages.domain.auth.password_hasher import (
password_validator,
)
from packages.domain.auth.session_store import RedisConfig, SessionStore
from packages.domain.auth.email_service import EmailConfig, EmailService
__all__ = [
"JWTService",
+43 -58
View File
@@ -2,55 +2,59 @@
JWT 工具类
提供 Token 签发、验证、刷新功能
"""
from datetime import datetime, timedelta
from typing import Dict, Any, Optional
from typing import Any, Dict, Optional
import jwt
from jwt.exceptions import ExpiredSignatureError, InvalidTokenError
class JWTConfig:
"""JWT 配置"""
# 从环境变量读取,这里先用默认值
SECRET_KEY: str = "your-secret-key-change-in-production"
ALGORITHM: str = "HS256"
ACCESS_TOKEN_EXPIRE_MINUTES: int = 30 # 30 分钟
REFRESH_TOKEN_EXPIRE_DAYS: int = 30 # 30 天
REFRESH_TOKEN_EXPIRE_DAYS: int = 30 # 30 天
class TokenType:
"""Token 类型"""
ACCESS = "access"
REFRESH = "refresh"
class JWTService:
"""JWT 服务类"""
def __init__(self, config: JWTConfig = None):
self.config = config or JWTConfig()
def create_access_token(
self,
user_id: str,
workspace_id: str,
role: str,
additional_claims: Optional[Dict[str, Any]] = None
additional_claims: Optional[Dict[str, Any]] = None,
) -> str:
"""
创建 access_token
Args:
user_id: 用户 ID
workspace_id: 工作空间 ID
role: 用户在该工作空间的角色
additional_claims: 额外的声明(可选)
Returns:
JWT Token 字符串
"""
now = datetime.utcnow()
expire = now + timedelta(minutes=self.config.ACCESS_TOKEN_EXPIRE_MINUTES)
payload = {
"sub": user_id, # subject (用户 ID)
"workspace_id": workspace_id,
@@ -59,34 +63,26 @@ class JWTService:
"iat": now, # issued at
"exp": expire, # expiration time
}
if additional_claims:
payload.update(additional_claims)
return jwt.encode(
payload,
self.config.SECRET_KEY,
algorithm=self.config.ALGORITHM
)
def create_refresh_token(
self,
user_id: str,
session_id: str
) -> str:
return jwt.encode(payload, self.config.SECRET_KEY, algorithm=self.config.ALGORITHM)
def create_refresh_token(self, user_id: str, session_id: str) -> str:
"""
创建 refresh_token
Args:
user_id: 用户 ID
session_id: Session ID(用于撤销)
Returns:
JWT Token 字符串
"""
now = datetime.utcnow()
expire = now + timedelta(days=self.config.REFRESH_TOKEN_EXPIRE_DAYS)
payload = {
"sub": user_id,
"session_id": session_id,
@@ -94,98 +90,87 @@ class JWTService:
"iat": now,
"exp": expire,
}
return jwt.encode(
payload,
self.config.SECRET_KEY,
algorithm=self.config.ALGORITHM
)
return jwt.encode(payload, self.config.SECRET_KEY, algorithm=self.config.ALGORITHM)
def verify_token(self, token: str) -> Dict[str, Any]:
"""
验证 Token 并解码
Args:
token: JWT Token 字符串
Returns:
Token payload
Raises:
ExpiredSignatureError: Token 已过期
InvalidTokenError: Token 无效
"""
try:
payload = jwt.decode(
token,
self.config.SECRET_KEY,
algorithms=[self.config.ALGORITHM]
)
payload = jwt.decode(token, self.config.SECRET_KEY, algorithms=[self.config.ALGORITHM])
return payload
except ExpiredSignatureError:
raise ExpiredSignatureError("Token has expired")
except InvalidTokenError as e:
raise InvalidTokenError(f"Invalid token: {str(e)}")
def verify_access_token(self, token: str) -> Dict[str, Any]:
"""
验证 access_token
Args:
token: JWT Token 字符串
Returns:
Token payload
Raises:
ValueError: Token 类型不是 access
ExpiredSignatureError: Token 已过期
InvalidTokenError: Token 无效
"""
payload = self.verify_token(token)
if payload.get("type") != TokenType.ACCESS:
raise ValueError("Token type must be 'access'")
return payload
def verify_refresh_token(self, token: str) -> Dict[str, Any]:
"""
验证 refresh_token
Args:
token: JWT Token 字符串
Returns:
Token payload
Raises:
ValueError: Token 类型不是 refresh
ExpiredSignatureError: Token 已过期
InvalidTokenError: Token 无效
"""
payload = self.verify_token(token)
if payload.get("type") != TokenType.REFRESH:
raise ValueError("Token type must be 'refresh'")
return payload
def decode_token_unsafe(self, token: str) -> Optional[Dict[str, Any]]:
"""
不验证签名地解码 Token(仅用于调试/日志)
Args:
token: JWT Token 字符串
Returns:
Token payload(如果解码失败返回 None)
"""
try:
return jwt.decode(
token,
options={"verify_signature": False}
)
return jwt.decode(token, options={"verify_signature": False})
except Exception:
return None
+40 -38
View File
@@ -2,97 +2,99 @@
密码哈希工具类
使用 bcrypt 安全存储密码
"""
import bcrypt
from typing import Optional
import bcrypt
class PasswordHasher:
"""密码哈希服务"""
def __init__(self, rounds: int = 12):
"""
初始化密码哈希器
Args:
rounds: bcrypt cost factor(默认 12,推荐范围 10-14
值越大越安全,但计算时间越长
"""
if rounds < 4 or rounds > 31:
raise ValueError("rounds must be between 4 and 31")
self.rounds = rounds
def hash_password(self, password: str) -> str:
"""
哈希密码
Args:
password: 明文密码
Returns:
bcrypt 哈希字符串(包含 salt
Raises:
ValueError: 密码为空
"""
if not password:
raise ValueError("Password cannot be empty")
# bcrypt 需要 bytes
password_bytes = password.encode('utf-8')
password_bytes = password.encode("utf-8")
# 生成 salt 并哈希
salt = bcrypt.gensalt(rounds=self.rounds)
hashed = bcrypt.hashpw(password_bytes, salt)
# 返回字符串(数据库存储)
return hashed.decode('utf-8')
return hashed.decode("utf-8")
def verify_password(self, password: str, hashed_password: str) -> bool:
"""
验证密码
Args:
password: 明文密码
hashed_password: 存储的哈希密码
Returns:
True 如果密码正确,否则 False
"""
if not password or not hashed_password:
return False
try:
password_bytes = password.encode('utf-8')
hashed_bytes = hashed_password.encode('utf-8')
password_bytes = password.encode("utf-8")
hashed_bytes = hashed_password.encode("utf-8")
return bcrypt.checkpw(password_bytes, hashed_bytes)
except Exception:
# 哈希格式错误或其他异常,返回 False
return False
def needs_rehash(self, hashed_password: str) -> bool:
"""
检查哈希是否需要重新计算
(当 cost factor 改变时需要重新哈希)
Args:
hashed_password: 存储的哈希密码
Returns:
True 如果需要重新哈希
"""
try:
hashed_bytes = hashed_password.encode('utf-8')
hashed_bytes = hashed_password.encode("utf-8")
current_rounds = bcrypt.getsalt(hashed_bytes)
# 提取当前的 cost factor
# bcrypt hash 格式: $2b$rounds$salt+hash
parts = hashed_password.split('$')
parts = hashed_password.split("$")
if len(parts) >= 3:
stored_rounds = int(parts[2])
return stored_rounds != self.rounds
return False
except Exception:
return False
@@ -100,7 +102,7 @@ class PasswordHasher:
class PasswordValidator:
"""密码强度验证器"""
def __init__(
self,
min_length: int = 8,
@@ -111,7 +113,7 @@ class PasswordValidator:
):
"""
初始化密码验证器
Args:
min_length: 最小长度
require_uppercase: 是否要求大写字母
@@ -124,37 +126,37 @@ class PasswordValidator:
self.require_lowercase = require_lowercase
self.require_digit = require_digit
self.require_special = require_special
def validate(self, password: str) -> tuple[bool, Optional[str]]:
"""
验证密码强度
Args:
password: 明文密码
Returns:
(是否有效, 错误信息)
"""
if not password:
return False, "Password cannot be empty"
if len(password) < self.min_length:
return False, f"Password must be at least {self.min_length} characters"
if self.require_uppercase and not any(c.isupper() for c in password):
return False, "Password must contain at least one uppercase letter"
if self.require_lowercase and not any(c.islower() for c in password):
return False, "Password must contain at least one lowercase letter"
if self.require_digit and not any(c.isdigit() for c in password):
return False, "Password must contain at least one digit"
if self.require_special:
special_chars = "!@#$%^&*()_+-=[]{}|;:,.<>?~"
if not any(c in special_chars for c in password):
return False, "Password must contain at least one special character"
return True, None
+1
View File
@@ -28,6 +28,7 @@ class ClassificationJobStatus(StrEnum):
class AssetClassification(StrEnum):
"""Asset classification categories."""
SCENIC = "scenic" # 风景
PRODUCT = "product" # 产品
PERSON = "person" # 人物
+4
View File
@@ -55,6 +55,7 @@ class Workspace:
class WorkspaceMemberRole(StrEnum):
"""工作空间成员角色"""
OWNER = "owner" # 所有者(创建者,唯一)
ADMIN = "admin" # 管理员(可管理成员和项目)
MEMBER = "member" # 成员(可创建和编辑项目)
@@ -63,6 +64,7 @@ class WorkspaceMemberRole(StrEnum):
class InvitationStatus(StrEnum):
"""邀请状态"""
PENDING = "pending" # 待处理
ACCEPTED = "accepted" # 已接受
DECLINED = "declined" # 已拒绝
@@ -72,6 +74,7 @@ class InvitationStatus(StrEnum):
@dataclass(slots=True)
class WorkspaceMember:
"""工作空间成员"""
id: str
workspace_id: str
user_id: str
@@ -83,6 +86,7 @@ class WorkspaceMember:
@dataclass(slots=True)
class WorkspaceInvitation:
"""工作空间邀请"""
id: str
workspace_id: str
inviter_user_id: str # 邀请人
+41 -39
View File
@@ -2,16 +2,18 @@
权限验证辅助函数
用于检查用户在工作空间中的权限
"""
from typing import Optional
from packages.domain.entities import WorkspaceMemberRole
class PermissionChecker:
"""权限检查器"""
def __init__(self, workspace_member_repository):
self.workspace_member_repository = workspace_member_repository
def check_workspace_access(
self,
workspace_id: str,
@@ -19,11 +21,11 @@ class PermissionChecker:
) -> tuple[bool, Optional[str]]:
"""
检查用户是否可以访问工作空间
Args:
workspace_id: 工作空间 ID
user_id: 用户 ID
Returns:
(是否有权限, 用户角色)
"""
@@ -31,12 +33,12 @@ class PermissionChecker:
workspace_id,
user_id,
)
if not member:
return False, None
return True, member.role
def check_is_owner(
self,
workspace_id: str,
@@ -44,11 +46,11 @@ class PermissionChecker:
) -> bool:
"""
检查用户是否是工作空间 Owner
Args:
workspace_id: 工作空间 ID
user_id: 用户 ID
Returns:
是否是 Owner
"""
@@ -56,9 +58,9 @@ class PermissionChecker:
workspace_id,
user_id,
)
return member is not None and member.role == WorkspaceMemberRole.OWNER
def check_is_admin_or_owner(
self,
workspace_id: str,
@@ -66,11 +68,11 @@ class PermissionChecker:
) -> bool:
"""
检查用户是否是工作空间 Admin 或 Owner
Args:
workspace_id: 工作空间 ID
user_id: 用户 ID
Returns:
是否是 Admin 或 Owner
"""
@@ -78,12 +80,12 @@ class PermissionChecker:
workspace_id,
user_id,
)
return member is not None and member.role in [
WorkspaceMemberRole.OWNER,
WorkspaceMemberRole.ADMIN,
]
def check_can_manage_members(
self,
workspace_id: str,
@@ -91,16 +93,16 @@ class PermissionChecker:
) -> bool:
"""
检查用户是否可以管理成员(邀请、移除、修改角色)
Args:
workspace_id: 工作空间 ID
user_id: 用户 ID
Returns:
是否可以管理成员
"""
return self.check_is_admin_or_owner(workspace_id, user_id)
def check_can_edit_workspace(
self,
workspace_id: str,
@@ -108,16 +110,16 @@ class PermissionChecker:
) -> bool:
"""
检查用户是否可以编辑工作空间(修改名称、设置等)
Args:
workspace_id: 工作空间 ID
user_id: 用户 ID
Returns:
是否可以编辑工作空间
"""
return self.check_is_admin_or_owner(workspace_id, user_id)
def check_can_create_project(
self,
workspace_id: str,
@@ -125,11 +127,11 @@ class PermissionChecker:
) -> bool:
"""
检查用户是否可以创建项目
Args:
workspace_id: 工作空间 ID
user_id: 用户 ID
Returns:
是否可以创建项目
"""
@@ -137,14 +139,14 @@ class PermissionChecker:
workspace_id,
user_id,
)
# Owner, Admin, Member 可以创建项目,Viewer 不可以
return member is not None and member.role in [
WorkspaceMemberRole.OWNER,
WorkspaceMemberRole.ADMIN,
WorkspaceMemberRole.MEMBER,
]
def check_can_edit_project(
self,
workspace_id: str,
@@ -152,17 +154,17 @@ class PermissionChecker:
) -> bool:
"""
检查用户是否可以编辑项目
Args:
workspace_id: 工作空间 ID
user_id: 用户 ID
Returns:
是否可以编辑项目
"""
# 与创建项目权限相同
return self.check_can_create_project(workspace_id, user_id)
def check_can_delete_project(
self,
workspace_id: str,
@@ -170,17 +172,17 @@ class PermissionChecker:
) -> bool:
"""
检查用户是否可以删除项目
Args:
workspace_id: 工作空间 ID
user_id: 用户 ID
Returns:
是否可以删除项目
"""
# 只有 Owner 和 Admin 可以删除项目
return self.check_is_admin_or_owner(workspace_id, user_id)
def check_can_view_workspace(
self,
workspace_id: str,
@@ -188,11 +190,11 @@ class PermissionChecker:
) -> bool:
"""
检查用户是否可以查看工作空间
Args:
workspace_id: 工作空间 ID
user_id: 用户 ID
Returns:
是否可以查看工作空间
"""
@@ -203,25 +205,25 @@ class PermissionChecker:
# 权限级别定义
class Permission:
"""权限常量"""
# 工作空间权限
WORKSPACE_VIEW = "workspace:view"
WORKSPACE_EDIT = "workspace:edit"
WORKSPACE_DELETE = "workspace:delete"
WORKSPACE_MANAGE_SUBSCRIPTION = "workspace:manage_subscription"
# 成员权限
MEMBER_VIEW = "member:view"
MEMBER_INVITE = "member:invite"
MEMBER_REMOVE = "member:remove"
MEMBER_UPDATE_ROLE = "member:update_role"
# 项目权限
PROJECT_VIEW = "project:view"
PROJECT_CREATE = "project:create"
PROJECT_EDIT = "project:edit"
PROJECT_DELETE = "project:delete"
# 资产权限
ASSET_VIEW = "asset:view"
ASSET_UPLOAD = "asset:upload"
@@ -289,11 +291,11 @@ ROLE_PERMISSIONS = {
def has_permission(role: str, permission: str) -> bool:
"""
检查角色是否有指定权限
Args:
role: 用户角色
permission: 权限标识
Returns:
是否有权限
"""
+11 -5
View File
@@ -1,4 +1,5 @@
"""项目管理领域对象:任务、里程碑、项目阶段"""
from __future__ import annotations
from dataclasses import dataclass, field
@@ -9,6 +10,7 @@ from uuid import uuid4
class TaskStatus(StrEnum):
"""任务状态"""
PENDING = "pending" # 待开始
IN_PROGRESS = "in_progress" # 进行中
BLOCKED = "blocked" # 阻塞
@@ -18,6 +20,7 @@ class TaskStatus(StrEnum):
class TaskPriority(StrEnum):
"""任务优先级"""
LOW = "low"
MEDIUM = "medium"
HIGH = "high"
@@ -27,6 +30,7 @@ class TaskPriority(StrEnum):
@dataclass(slots=True)
class Task:
"""任务实体"""
id: str
project_id: str
workspace_id: str
@@ -66,7 +70,7 @@ class Task:
raise ValueError("project_id 不能为空")
if not workspace_id.strip():
raise ValueError("workspace_id 不能为空")
return cls(
id=uuid4().hex,
project_id=project_id.strip(),
@@ -84,7 +88,7 @@ class Task:
"""更新任务状态"""
self.status = new_status
self.updated_at = datetime.now(timezone.utc)
# 自动设置实际开始/结束时间
if new_status == TaskStatus.IN_PROGRESS and self.actual_start_date is None:
self.actual_start_date = datetime.now(timezone.utc)
@@ -98,7 +102,7 @@ class Task:
raise ValueError("进度必须在 0-100 之间")
self.progress = progress
self.updated_at = datetime.now(timezone.utc)
# 自动更新状态
if progress > 0 and self.status == TaskStatus.PENDING:
self.status = TaskStatus.IN_PROGRESS
@@ -127,6 +131,7 @@ class Task:
@dataclass(slots=True)
class Milestone:
"""里程碑实体"""
id: str
project_id: str
workspace_id: str
@@ -155,7 +160,7 @@ class Milestone:
raise ValueError("project_id 不能为空")
if not workspace_id.strip():
raise ValueError("workspace_id 不能为空")
return cls(
id=uuid4().hex,
project_id=project_id.strip(),
@@ -183,6 +188,7 @@ class Milestone:
@dataclass(slots=True)
class TaskIssue:
"""任务问题/卡点实体"""
id: str
task_id: str
project_id: str
@@ -215,7 +221,7 @@ class TaskIssue:
raise ValueError("project_id 不能为空")
if not workspace_id.strip():
raise ValueError("workspace_id 不能为空")
return cls(
id=uuid4().hex,
task_id=task_id.strip(),
+41 -38
View File
@@ -2,12 +2,13 @@
配额检查服务
用于检查工作空间是否超出配额限制
"""
from typing import Optional
class QuotaChecker:
"""配额检查器"""
def __init__(
self,
workspace_repository,
@@ -15,33 +16,36 @@ class QuotaChecker:
):
self.workspace_repository = workspace_repository
self.project_repository = project_repository
def check_can_create_project(
self,
workspace_id: str,
) -> tuple[bool, Optional[str]]:
"""
检查是否可以创建项目
Args:
workspace_id: 工作空间 ID
Returns:
(是否可以, 错误信息)
"""
workspace = self.workspace_repository.find_by_id(workspace_id)
if not workspace:
return False, "Workspace not found"
# 获取当前项目数量
current_count = self.project_repository.count_by_workspace(workspace_id)
# 检查是否超出配额(999999 表示无限)
if workspace.max_projects != 999999 and current_count >= workspace.max_projects:
return False, f"Project limit reached ({workspace.max_projects}). Upgrade your plan to create more projects."
return (
False,
f"Project limit reached ({workspace.max_projects}). Upgrade your plan to create more projects.",
)
return True, None
def check_storage_available(
self,
workspace_id: str,
@@ -49,57 +53,56 @@ class QuotaChecker:
) -> tuple[bool, Optional[str]]:
"""
检查存储空间是否足够
Args:
workspace_id: 工作空间 ID
additional_gb: 需要的额外存储空间(GB)
Returns:
(是否可以, 错误信息)
"""
workspace = self.workspace_repository.find_by_id(workspace_id)
if not workspace:
return False, "Workspace not found"
# 检查存储空间
new_usage = workspace.used_storage_gb + additional_gb
if new_usage > workspace.max_storage_gb:
remaining = workspace.max_storage_gb - workspace.used_storage_gb
return False, f"Storage limit exceeded. Available: {remaining:.2f}GB, Required: {additional_gb:.2f}GB. Upgrade your plan for more storage."
return (
False,
f"Storage limit exceeded. Available: {remaining:.2f}GB, Required: {additional_gb:.2f}GB. Upgrade your plan for more storage.",
)
return True, None
def get_quota_status(self, workspace_id: str) -> dict:
"""
获取配额使用状态
Args:
workspace_id: 工作空间 ID
Returns:
配额状态信息
"""
workspace = self.workspace_repository.find_by_id(workspace_id)
if not workspace:
return None
# 获取项目数量
project_count = self.project_repository.count_by_workspace(workspace_id)
# 计算使用率
project_usage_percent = (
(project_count / workspace.max_projects * 100)
if workspace.max_projects != 999999
else 0 # 无限制
(project_count / workspace.max_projects * 100) if workspace.max_projects != 999999 else 0 # 无限制
)
storage_usage_percent = (
(workspace.used_storage_gb / workspace.max_storage_gb * 100)
if workspace.max_storage_gb > 0
else 0
(workspace.used_storage_gb / workspace.max_storage_gb * 100) if workspace.max_storage_gb > 0 else 0
)
return {
"workspace_id": workspace.id,
"subscription_plan": workspace.subscription_plan,
@@ -116,7 +119,7 @@ class QuotaChecker:
"usage_percent": storage_usage_percent,
},
}
def update_storage_usage(
self,
workspace_id: str,
@@ -124,34 +127,34 @@ class QuotaChecker:
) -> tuple[bool, Optional[str]]:
"""
更新存储使用量
Args:
workspace_id: 工作空间 ID
delta_gb: 变化量(正数为增加,负数为减少)
Returns:
(是否成功, 错误信息)
"""
workspace = self.workspace_repository.find_by_id(workspace_id)
if not workspace:
return False, "Workspace not found"
# 更新使用量
new_usage = workspace.used_storage_gb + delta_gb
# 不能为负数
if new_usage < 0:
new_usage = 0
workspace.used_storage_gb = new_usage
self.workspace_repository.save(workspace)
return True, None
class QuotaWarningLevel:
"""配额警告级别"""
NORMAL = "normal" # <80%
WARNING = "warning" # 80-90%
CRITICAL = "critical" # 90-100%
@@ -161,10 +164,10 @@ class QuotaWarningLevel:
def get_warning_level(usage_percent: float) -> str:
"""
根据使用率获取警告级别
Args:
usage_percent: 使用率(0-100
Returns:
警告级别
"""
+5 -1
View File
@@ -3,7 +3,11 @@
from .asset_library_repository import AssetLibraryRepository
from .asset_repository import AssetRepository
from .ingest_job_repository import IngestJobRepository
from .project_management_repositories import MilestoneRepository, TaskIssueRepository, TaskRepository
from .project_management_repositories import (
MilestoneRepository,
TaskIssueRepository,
TaskRepository,
)
from .project_repository import ProjectRepository
__all__ = [
+4 -8
View File
@@ -6,14 +6,10 @@ from packages.domain import GeneratedVideo
class GeneratedVideoRepository(Protocol):
def create(self, video: GeneratedVideo) -> GeneratedVideo:
...
def create(self, video: GeneratedVideo) -> GeneratedVideo: ...
def get(self, video_id: str) -> GeneratedVideo | None:
...
def get(self, video_id: str) -> GeneratedVideo | None: ...
def list_by_project(self, project_id: str) -> list[GeneratedVideo]:
...
def list_by_project(self, project_id: str) -> list[GeneratedVideo]: ...
def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]:
...
def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]: ...
+4 -8
View File
@@ -6,14 +6,10 @@ from packages.domain import GenerationTask
class GenerationTaskRepository(Protocol):
def create(self, task: GenerationTask) -> GenerationTask:
...
def create(self, task: GenerationTask) -> GenerationTask: ...
def get(self, task_id: str) -> GenerationTask | None:
...
def get(self, task_id: str) -> GenerationTask | None: ...
def list_by_project(self, project_id: str) -> list[GenerationTask]:
...
def list_by_project(self, project_id: str) -> list[GenerationTask]: ...
def update(self, task: GenerationTask) -> GenerationTask:
...
def update(self, task: GenerationTask) -> GenerationTask: ...
@@ -1,4 +1,5 @@
"""项目管理 Repository 接口定义"""
from abc import ABC, abstractmethod
from packages.domain import Milestone, Task, TaskIssue
+6 -4
View File
@@ -1,29 +1,31 @@
"""
Project 仓储接口
"""
from abc import ABC, abstractmethod
from typing import Optional
from packages.domain.entities import Project
class ProjectRepository(ABC):
"""Project 仓储接口"""
@abstractmethod
def save(self, project: Project) -> None:
"""保存项目"""
pass
@abstractmethod
def find_by_id(self, project_id: str) -> Optional[Project]:
"""根据 ID 查找项目"""
pass
@abstractmethod
def count_by_workspace(self, workspace_id: str) -> int:
"""统计工作空间的项目数量"""
pass
@abstractmethod
def delete(self, project_id: str) -> bool:
"""删除项目"""
+9 -7
View File
@@ -1,44 +1,46 @@
"""
用户仓储接口
"""
from abc import ABC, abstractmethod
from typing import Optional
from packages.domain.entities import User
class UserRepository(ABC):
"""用户仓储接口"""
@abstractmethod
def save(self, user: User) -> None:
"""保存用户"""
pass
@abstractmethod
def find_by_id(self, user_id: str) -> Optional[User]:
"""根据 ID 查找用户"""
pass
@abstractmethod
def find_by_email(self, email: str) -> Optional[User]:
"""根据邮箱查找用户"""
pass
@abstractmethod
def find_by_username(self, username: str) -> Optional[User]:
"""根据用户名查找用户"""
pass
@abstractmethod
def find_by_verification_token(self, token: str) -> Optional[User]:
"""根据邮箱验证令牌查找用户"""
pass
@abstractmethod
def find_by_password_reset_token(self, token: str) -> Optional[User]:
"""根据密码重置令牌查找用户"""
pass
@abstractmethod
def delete(self, user_id: str) -> bool:
"""删除用户"""
@@ -1,29 +1,31 @@
"""
WorkspaceInvitation 仓储接口
"""
from abc import ABC, abstractmethod
from typing import Optional
from packages.domain.entities import WorkspaceInvitation
class WorkspaceInvitationRepository(ABC):
"""WorkspaceInvitation 仓储接口"""
@abstractmethod
def save(self, invitation: WorkspaceInvitation) -> None:
"""保存邀请"""
pass
@abstractmethod
def find_by_id(self, invitation_id: str) -> Optional[WorkspaceInvitation]:
"""根据 ID 查找邀请"""
pass
@abstractmethod
def find_by_token(self, token: str) -> Optional[WorkspaceInvitation]:
"""根据令牌查找邀请"""
pass
@abstractmethod
def find_pending_by_workspace_and_email(
self,
@@ -32,7 +34,7 @@ class WorkspaceInvitationRepository(ABC):
) -> Optional[WorkspaceInvitation]:
"""查找 workspace 和邮箱的待处理邀请"""
pass
@abstractmethod
def delete(self, invitation_id: str) -> bool:
"""删除邀请"""
+10 -8
View File
@@ -1,24 +1,26 @@
"""
WorkspaceMember 仓储接口
"""
from abc import ABC, abstractmethod
from typing import Optional, List
from typing import List, Optional
from packages.domain.entities import WorkspaceMember
class WorkspaceMemberRepository(ABC):
"""WorkspaceMember 仓储接口"""
@abstractmethod
def save(self, member: WorkspaceMember) -> None:
"""保存成员"""
pass
@abstractmethod
def find_by_id(self, member_id: str) -> Optional[WorkspaceMember]:
"""根据 ID 查找成员"""
pass
@abstractmethod
def find_by_workspace_and_user(
self,
@@ -27,22 +29,22 @@ class WorkspaceMemberRepository(ABC):
) -> Optional[WorkspaceMember]:
"""根据 workspace 和 user 查找成员"""
pass
@abstractmethod
def find_by_user(self, user_id: str) -> List[WorkspaceMember]:
"""查找用户的所有成员记录"""
pass
@abstractmethod
def find_by_workspace(self, workspace_id: str) -> List[WorkspaceMember]:
"""查找 workspace 的所有成员"""
pass
@abstractmethod
def count_by_workspace(self, workspace_id: str) -> int:
"""统计 workspace 的成员数量"""
pass
@abstractmethod
def delete(self, member_id: str) -> bool:
"""删除成员"""

Some files were not shown because too many files have changed in this diff Show More