diff --git a/alembic/env.py b/alembic/env.py index f48af7da6..b2c4797dd 100644 --- a/alembic/env.py +++ b/alembic/env.py @@ -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 diff --git a/alembic/versions/001_current_schema_baseline.py b/alembic/versions/001_current_schema_baseline.py index 305981f85..aadc682fa 100644 --- a/alembic/versions/001_current_schema_baseline.py +++ b/alembic/versions/001_current_schema_baseline.py @@ -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: diff --git a/apps/api/app/api/router.py b/apps/api/app/api/router.py index ba8883543..2a93c42a1 100644 --- a/apps/api/app/api/router.py +++ b/apps/api/app/api/router.py @@ -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() diff --git a/apps/api/app/api/routes/asset_libraries.py b/apps/api/app/api/routes/asset_libraries.py index 54018327b..d6e6e46eb 100644 --- a/apps/api/app/api/routes/asset_libraries.py +++ b/apps/api/app/api/routes/asset_libraries.py @@ -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() diff --git a/apps/api/app/api/routes/assets.py b/apps/api/app/api/routes/assets.py index 1ab21067d..cca253167 100644 --- a/apps/api/app/api/routes/assets.py +++ b/apps/api/app/api/routes/assets.py @@ -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() diff --git a/apps/api/app/api/routes/auth.py b/apps/api/app/api/routes/auth.py index 4fdab1b83..b3d53ccd0 100644 --- a/apps/api/app/api/routes/auth.py +++ b/apps/api/app/api/routes/auth.py @@ -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"} diff --git a/apps/api/app/api/routes/auth_simple.py b/apps/api/app/api/routes/auth_simple.py index 49b0bf082..e779e9ca5 100644 --- a/apps/api/app/api/routes/auth_simple.py +++ b/apps/api/app/api/routes/auth_simple.py @@ -1,17 +1,18 @@ """ 认证 API(SQLAlchemy ORM) """ -from datetime import datetime, timedelta, timezone + import hashlib import secrets +from datetime import datetime, timedelta, timezone import jwt +from app.config import settings +from app.dependencies import get_db_session from fastapi import APIRouter, Depends, HTTPException, status from pydantic import BaseModel, EmailStr from sqlalchemy.orm import Session -from app.config import settings -from app.dependencies import get_db_session from packages.adapters.sqlalchemy_impl.models import UserModel from packages.domain.auth import password_hasher, password_validator @@ -169,4 +170,7 @@ async def login(request: LoginRequest, db: Session = Depends(get_db_session)): @router.get("/me") async def get_current_user_info(): - raise HTTPException(status_code=status.HTTP_501_NOT_IMPLEMENTED, detail="/auth/me requires bearer-token dependency integration") + raise HTTPException( + status_code=status.HTTP_501_NOT_IMPLEMENTED, + detail="/auth/me requires bearer-token dependency integration", + ) diff --git a/apps/api/app/api/routes/classification_jobs.py b/apps/api/app/api/routes/classification_jobs.py index f69172ea3..e5161bd58 100644 --- a/apps/api/app/api/routes/classification_jobs.py +++ b/apps/api/app/api/routes/classification_jobs.py @@ -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() diff --git a/apps/api/app/api/routes/generated_videos.py b/apps/api/app/api/routes/generated_videos.py index ba03d3318..869db4bfc 100644 --- a/apps/api/app/api/routes/generated_videos.py +++ b/apps/api/app/api/routes/generated_videos.py @@ -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, diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index 1da9ca981..24aeca7f4 100644 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -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, diff --git a/apps/api/app/api/routes/health.py b/apps/api/app/api/routes/health.py index 64c97b7f8..bb85ff5c7 100644 --- a/apps/api/app/api/routes/health.py +++ b/apps/api/app/api/routes/health.py @@ -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: diff --git a/apps/api/app/api/routes/ingest_jobs.py b/apps/api/app/api/routes/ingest_jobs.py index 1737c6853..1c6421055 100644 --- a/apps/api/app/api/routes/ingest_jobs.py +++ b/apps/api/app/api/routes/ingest_jobs.py @@ -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() diff --git a/apps/api/app/api/routes/project_management.py b/apps/api/app/api/routes/project_management.py index 5228a4f3f..ee488b905 100644 --- a/apps/api/app/api/routes/project_management.py +++ b/apps/api/app/api/routes/project_management.py @@ -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() diff --git a/apps/api/app/api/routes/projects.py b/apps/api/app/api/routes/projects.py index ba1ee7c2c..7ac381ff3 100644 --- a/apps/api/app/api/routes/projects.py +++ b/apps/api/app/api/routes/projects.py @@ -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() diff --git a/apps/api/app/api/routes/upload.py b/apps/api/app/api/routes/upload.py index 9736e4069..ff031338a 100644 --- a/apps/api/app/api/routes/upload.py +++ b/apps/api/app/api/routes/upload.py @@ -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() diff --git a/apps/api/app/api/routes/workspaces.py b/apps/api/app/api/routes/workspaces.py index 5d147e63e..12025a6ed 100644 --- a/apps/api/app/api/routes/workspaces.py +++ b/apps/api/app/api/routes/workspaces.py @@ -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"} diff --git a/apps/api/app/config.py b/apps/api/app/config.py index 1450e6b9b..f8fd6889b 100644 --- a/apps/api/app/config.py +++ b/apps/api/app/config.py @@ -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): diff --git a/apps/api/app/core/celery_app.py b/apps/api/app/core/celery_app.py index c6b9d215d..52b515335 100644 --- a/apps/api/app/core/celery_app.py +++ b/apps/api/app/core/celery_app.py @@ -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") diff --git a/apps/api/app/core/database.py b/apps/api/app/core/database.py index 5f6422aa2..7caa3821c 100644 --- a/apps/api/app/core/database.py +++ b/apps/api/app/core/database.py @@ -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): diff --git a/apps/api/app/core/storage.py b/apps/api/app/core/storage.py index 149b96bb0..aeb37ca1a 100644 --- a/apps/api/app/core/storage.py +++ b/apps/api/app/core/storage.py @@ -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: 是否存在 """ diff --git a/apps/api/app/db.py b/apps/api/app/db.py index 74b1dd937..70c689f75 100644 --- a/apps/api/app/db.py +++ b/apps/api/app/db.py @@ -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( diff --git a/apps/api/app/dependencies.py b/apps/api/app/dependencies.py index 7ca4e8584..c0cf958d2 100644 --- a/apps/api/app/dependencies.py +++ b/apps/api/app/dependencies.py @@ -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) diff --git a/apps/api/app/middleware/auth.py b/apps/api/app/middleware/auth.py index 31aee583a..d4fd33fcf 100644 --- a/apps/api/app/middleware/auth.py +++ b/apps/api/app/middleware/auth.py @@ -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 diff --git a/apps/api/app/middleware/exceptions.py b/apps/api/app/middleware/exceptions.py index ef3c76a00..226ef2a30 100644 --- a/apps/api/app/middleware/exceptions.py +++ b/apps/api/app/middleware/exceptions.py @@ -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, diff --git a/apps/api/app/middleware/logging.py b/apps/api/app/middleware/logging.py index 324958ad4..deded104f 100644 --- a/apps/api/app/middleware/logging.py +++ b/apps/api/app/middleware/logging.py @@ -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 diff --git a/apps/api/app/middleware/monitoring.py b/apps/api/app/middleware/monitoring.py index fb62c3867..77ad2deae 100644 --- a/apps/api/app/middleware/monitoring.py +++ b/apps/api/app/middleware/monitoring.py @@ -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 { diff --git a/apps/api/app/middleware/versioning.py b/apps/api/app/middleware/versioning.py index afb40e091..0055aed43 100644 --- a/apps/api/app/middleware/versioning.py +++ b/apps/api/app/middleware/versioning.py @@ -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("/") diff --git a/apps/api/app/schemas/__init__.py b/apps/api/app/schemas/__init__.py index 3bf648c71..ee65ea2f6 100644 --- a/apps/api/app/schemas/__init__.py +++ b/apps/api/app/schemas/__init__.py @@ -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 diff --git a/apps/api/main.py b/apps/api/main.py index 148477f2e..73a2302bb 100644 --- a/apps/api/main.py +++ b/apps/api/main.py @@ -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", diff --git a/apps/worker/tasks.py b/apps/worker/tasks.py index ad02ba73a..df4461a51 100644 --- a/apps/worker/tasks.py +++ b/apps/worker/tasks.py @@ -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 diff --git a/apps/worker/video_processing/__init__.py b/apps/worker/video_processing/__init__.py index 0bb54f745..30465a107 100644 --- a/apps/worker/video_processing/__init__.py +++ b/apps/worker/video_processing/__init__.py @@ -1,6 +1,7 @@ """ 视频处理模块 """ + from .processor import VideoProcessor, VideoResult __all__ = ["VideoProcessor", "VideoResult"] diff --git a/apps/worker/video_processing/processor.py b/apps/worker/video_processing/processor.py index c5c203fb5..7af9dbf6b 100644 --- a/apps/worker/video_processing/processor.py +++ b/apps/worker/video_processing/processor.py @@ -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 diff --git a/apps/worker/worker_app/celery_app.py b/apps/worker/worker_app/celery_app.py index d352abc01..a14dd2e29 100644 --- a/apps/worker/worker_app/celery_app.py +++ b/apps/worker/worker_app/celery_app.py @@ -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 diff --git a/apps/worker/worker_app/core/config.py b/apps/worker/worker_app/core/config.py index 4e5658c3c..b8d23ef13 100644 --- a/apps/worker/worker_app/core/config.py +++ b/apps/worker/worker_app/core/config.py @@ -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): diff --git a/apps/worker/worker_app/db.py b/apps/worker/worker_app/db.py index 260f57af3..f051e0f23 100644 --- a/apps/worker/worker_app/db.py +++ b/apps/worker/worker_app/db.py @@ -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) diff --git a/apps/worker/worker_app/tasks/classification.py b/apps/worker/worker_app/tasks/classification.py index e9c15cae8..7d00ce342 100644 --- a/apps/worker/worker_app/tasks/classification.py +++ b/apps/worker/worker_app/tasks/classification.py @@ -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, diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 74d605360..128ede28b 100644 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -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) diff --git a/apps/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py index a49bf3e22..cfdd5e6b2 100644 --- a/apps/worker/worker_app/tasks/ingest.py +++ b/apps/worker/worker_app/tasks/ingest.py @@ -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 diff --git a/packages/adapters/in_memory/asset_library_repository.py b/packages/adapters/in_memory/asset_library_repository.py index fd086a05c..1c218e805 100644 --- a/packages/adapters/in_memory/asset_library_repository.py +++ b/packages/adapters/in_memory/asset_library_repository.py @@ -1,4 +1,5 @@ """AssetLibrary InMemory Repository 实现""" + from packages.domain import AssetLibrary, AssetLibraryKind diff --git a/packages/adapters/in_memory/asset_repository.py b/packages/adapters/in_memory/asset_repository.py index cfd9cdaab..c275f3ebd 100644 --- a/packages/adapters/in_memory/asset_repository.py +++ b/packages/adapters/in_memory/asset_repository.py @@ -1,4 +1,5 @@ """Asset InMemory Repository 实现""" + from packages.domain import Asset diff --git a/packages/adapters/in_memory/project_management_repositories.py b/packages/adapters/in_memory/project_management_repositories.py index 44257d050..2f1132af0 100644 --- a/packages/adapters/in_memory/project_management_repositories.py +++ b/packages/adapters/in_memory/project_management_repositories.py @@ -1,4 +1,5 @@ """项目管理 In-Memory Repository 实现""" + from packages.domain import Milestone, Task, TaskIssue from packages.ports.project_management_repositories import ( MilestoneRepository, diff --git a/packages/adapters/in_memory/user_repository.py b/packages/adapters/in_memory/user_repository.py index f93eb0299..7eebc680f 100644 --- a/packages/adapters/in_memory/user_repository.py +++ b/packages/adapters/in_memory/user_repository.py @@ -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 diff --git a/packages/adapters/in_memory/workspace_invitation_repository.py b/packages/adapters/in_memory/workspace_invitation_repository.py index bf0a8bde6..c94433066 100644 --- a/packages/adapters/in_memory/workspace_invitation_repository.py +++ b/packages/adapters/in_memory/workspace_invitation_repository.py @@ -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 diff --git a/packages/adapters/in_memory/workspace_member_repository.py b/packages/adapters/in_memory/workspace_member_repository.py index 8694792dc..a0f83a161 100644 --- a/packages/adapters/in_memory/workspace_member_repository.py +++ b/packages/adapters/in_memory/workspace_member_repository.py @@ -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 diff --git a/packages/adapters/in_memory/workspace_repository.py b/packages/adapters/in_memory/workspace_repository.py index 5d2fbe439..3c3f6606e 100644 --- a/packages/adapters/in_memory/workspace_repository.py +++ b/packages/adapters/in_memory/workspace_repository.py @@ -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: diff --git a/packages/adapters/postgres/__init__.py b/packages/adapters/postgres/__init__.py index bcad4c3c4..be8868995 100644 --- a/packages/adapters/postgres/__init__.py +++ b/packages/adapters/postgres/__init__.py @@ -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") diff --git a/packages/adapters/postgres/asset_repository.py b/packages/adapters/postgres/asset_repository.py index 1e8f2eab2..503905b33 100644 --- a/packages/adapters/postgres/asset_repository.py +++ b/packages/adapters/postgres/asset_repository.py @@ -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 diff --git a/packages/adapters/postgres/connection_pool.py b/packages/adapters/postgres/connection_pool.py index 46d08ee6c..30a41b1ae 100644 --- a/packages/adapters/postgres/connection_pool.py +++ b/packages/adapters/postgres/connection_pool.py @@ -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: diff --git a/packages/adapters/postgres/project_repository.py b/packages/adapters/postgres/project_repository.py index a1c8f70e8..994333fc2 100644 --- a/packages/adapters/postgres/project_repository.py +++ b/packages/adapters/postgres/project_repository.py @@ -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( diff --git a/packages/adapters/postgres/user_repository.py b/packages/adapters/postgres/user_repository.py index 9bb448e60..b165393da 100644 --- a/packages/adapters/postgres/user_repository.py +++ b/packages/adapters/postgres/user_repository.py @@ -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( diff --git a/packages/adapters/postgres/workspace_invitation_repository.py b/packages/adapters/postgres/workspace_invitation_repository.py index 0e1c0b464..9107ee60c 100644 --- a/packages/adapters/postgres/workspace_invitation_repository.py +++ b/packages/adapters/postgres/workspace_invitation_repository.py @@ -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( diff --git a/packages/adapters/postgres/workspace_member_repository.py b/packages/adapters/postgres/workspace_member_repository.py index 3da29049d..85f82744c 100644 --- a/packages/adapters/postgres/workspace_member_repository.py +++ b/packages/adapters/postgres/workspace_member_repository.py @@ -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( diff --git a/packages/adapters/postgres/workspace_repository.py b/packages/adapters/postgres/workspace_repository.py index f69361d94..775a9b794 100644 --- a/packages/adapters/postgres/workspace_repository.py +++ b/packages/adapters/postgres/workspace_repository.py @@ -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( diff --git a/packages/adapters/redis/__init__.py b/packages/adapters/redis/__init__.py index 1d7a1be95..c7e4ebb75 100644 --- a/packages/adapters/redis/__init__.py +++ b/packages/adapters/redis/__init__.py @@ -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"] diff --git a/packages/adapters/redis/session_store.py b/packages/adapters/redis/session_store.py index 574537107..c09ea413b 100644 --- a/packages/adapters/redis/session_store.py +++ b/packages/adapters/redis/session_store.py @@ -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: 是否存在 """ diff --git a/packages/adapters/smtp/__init__.py b/packages/adapters/smtp/__init__.py index 63833654e..00dca6fca 100644 --- a/packages/adapters/smtp/__init__.py +++ b/packages/adapters/smtp/__init__.py @@ -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"] diff --git a/packages/adapters/smtp/email_service.py b/packages/adapters/smtp/email_service.py index c64d75cbc..835c6bbcd 100644 --- a/packages/adapters/smtp/email_service.py +++ b/packages/adapters/smtp/email_service.py @@ -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""" @@ -154,7 +152,7 @@ class EmailService: """ - + 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""" @@ -223,7 +221,7 @@ class EmailService: """ - + 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""" @@ -307,7 +305,7 @@ class EmailService: """ - + text_body = f""" 工作空间邀请 @@ -321,7 +319,7 @@ class EmailService: 如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。 """ - + return self.send_email(to_email, subject, html_body, text_body) diff --git a/packages/adapters/sqlalchemy_impl/__init__.py b/packages/adapters/sqlalchemy_impl/__init__.py index 592ed3281..12a24f279 100644 --- a/packages/adapters/sqlalchemy_impl/__init__.py +++ b/packages/adapters/sqlalchemy_impl/__init__.py @@ -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", diff --git a/packages/adapters/sqlalchemy_impl/asset_repository.py b/packages/adapters/sqlalchemy_impl/asset_repository.py index d3aa5707d..c08854703 100644 --- a/packages/adapters/sqlalchemy_impl/asset_repository.py +++ b/packages/adapters/sqlalchemy_impl/asset_repository.py @@ -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, diff --git a/packages/adapters/sqlalchemy_impl/generated_video_repository.py b/packages/adapters/sqlalchemy_impl/generated_video_repository.py index 752244320..376da70bd 100644 --- a/packages/adapters/sqlalchemy_impl/generated_video_repository.py +++ b/packages/adapters/sqlalchemy_impl/generated_video_repository.py @@ -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] diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index be424eca0..4b19b3e17 100644 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -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() diff --git a/packages/adapters/sqlalchemy_impl/project_management_repositories.py b/packages/adapters/sqlalchemy_impl/project_management_repositories.py index 3790234b4..0407b9eee 100644 --- a/packages/adapters/sqlalchemy_impl/project_management_repositories.py +++ b/packages/adapters/sqlalchemy_impl/project_management_repositories.py @@ -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 diff --git a/packages/adapters/sqlalchemy_impl/session.py b/packages/adapters/sqlalchemy_impl/session.py index 4f1057574..c48d33d72 100644 --- a/packages/adapters/sqlalchemy_impl/session.py +++ b/packages/adapters/sqlalchemy_impl/session.py @@ -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() diff --git a/packages/adapters/sqlite_tracker/__init__.py b/packages/adapters/sqlite_tracker/__init__.py index fdfba5bef..9f17f75ac 100644 --- a/packages/adapters/sqlite_tracker/__init__.py +++ b/packages/adapters/sqlite_tracker/__init__.py @@ -1,8 +1,9 @@ """SQLite Tracker Adapter""" + from .project_management_repositories import ( - SQLiteTaskRepository, SQLiteMilestoneRepository, SQLiteTaskIssueRepository, + SQLiteTaskRepository, ) __all__ = [ diff --git a/packages/adapters/sqlite_tracker/project_management_repositories.py b/packages/adapters/sqlite_tracker/project_management_repositories.py index 2079b0b27..ff4214d67 100644 --- a/packages/adapters/sqlite_tracker/project_management_repositories.py +++ b/packages/adapters/sqlite_tracker/project_management_repositories.py @@ -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 diff --git a/packages/application/__init__.py b/packages/application/__init__.py index 8ed50eea9..c8f35cb83 100644 --- a/packages/application/__init__.py +++ b/packages/application/__init__.py @@ -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 diff --git a/packages/application/auth/__init__.py b/packages/application/auth/__init__.py index ea8e45585..cd976fe79 100644 --- a/packages/application/auth/__init__.py +++ b/packages/application/auth/__init__.py @@ -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__ = [ diff --git a/packages/application/auth/login_use_case.py b/packages/application/auth/login_use_case.py index 484c4f979..febd925b4 100644 --- a/packages/application/auth/login_use_case.py +++ b/packages/application/auth/login_use_case.py @@ -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)}" diff --git a/packages/application/auth/password_reset_use_case.py b/packages/application/auth/password_reset_use_case.py index b482e4c8c..f43b9e17b 100644 --- a/packages/application/auth/password_reset_use_case.py +++ b/packages/application/auth/password_reset_use_case.py @@ -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)}" diff --git a/packages/application/auth/register_user_use_case.py b/packages/application/auth/register_user_use_case.py index be43cd276..a451fca72 100644 --- a/packages/application/auth/register_user_use_case.py +++ b/packages/application/auth/register_user_use_case.py @@ -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)}" diff --git a/packages/application/common/pagination.py b/packages/application/common/pagination.py index f700fb689..0705be447 100644 --- a/packages/application/common/pagination.py +++ b/packages/application/common/pagination.py @@ -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, diff --git a/packages/application/get_task_detail_use_case.py b/packages/application/get_task_detail_use_case.py index 89c7bafcb..f7ecf370f 100644 --- a/packages/application/get_task_detail_use_case.py +++ b/packages/application/get_task_detail_use_case.py @@ -1,4 +1,5 @@ """获取单个任务详情用例""" + from packages.domain import Task from packages.ports import TaskRepository diff --git a/packages/application/project_management_use_cases.py b/packages/application/project_management_use_cases.py index 9efdba7f6..64dcf05cc 100644 --- a/packages/application/project_management_use_cases.py +++ b/packages/application/project_management_use_cases.py @@ -1,4 +1,5 @@ """项目管理 Use Cases""" + from packages.domain import Milestone, Task, TaskIssue, TaskPriority, TaskStatus from packages.ports import MilestoneRepository, TaskIssueRepository, TaskRepository diff --git a/packages/application/update_task_use_case.py b/packages/application/update_task_use_case.py index 793edb954..7c742284c 100644 --- a/packages/application/update_task_use_case.py +++ b/packages/application/update_task_use_case.py @@ -1,4 +1,5 @@ """更新任务基本信息用例""" + from packages.domain import Task from packages.ports import TaskRepository diff --git a/packages/application/workspace/__init__.py b/packages/application/workspace/__init__.py index b255ead91..165ad98bd 100644 --- a/packages/application/workspace/__init__.py +++ b/packages/application/workspace/__init__.py @@ -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__ = [ diff --git a/packages/application/workspace/accept_invitation_use_case.py b/packages/application/workspace/accept_invitation_use_case.py index 7cbc0a061..5450d9da3 100644 --- a/packages/application/workspace/accept_invitation_use_case.py +++ b/packages/application/workspace/accept_invitation_use_case.py @@ -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)}" diff --git a/packages/application/workspace/create_workspace_use_case.py b/packages/application/workspace/create_workspace_use_case.py index 678e219d5..1a87fca8e 100644 --- a/packages/application/workspace/create_workspace_use_case.py +++ b/packages/application/workspace/create_workspace_use_case.py @@ -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)}" diff --git a/packages/application/workspace/invite_member_use_case.py b/packages/application/workspace/invite_member_use_case.py index c5b9ab160..d907cedc2 100644 --- a/packages/application/workspace/invite_member_use_case.py +++ b/packages/application/workspace/invite_member_use_case.py @@ -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)}" diff --git a/packages/application/workspace/list_members_use_case.py b/packages/application/workspace/list_members_use_case.py index eda1913f2..afe4a42c5 100644 --- a/packages/application/workspace/list_members_use_case.py +++ b/packages/application/workspace/list_members_use_case.py @@ -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)}" diff --git a/packages/application/workspace/list_workspaces_use_case.py b/packages/application/workspace/list_workspaces_use_case.py index 048f89117..1c27e3841 100644 --- a/packages/application/workspace/list_workspaces_use_case.py +++ b/packages/application/workspace/list_workspaces_use_case.py @@ -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)}" diff --git a/packages/application/workspace/remove_member_use_case.py b/packages/application/workspace/remove_member_use_case.py index d23d09192..046cf5574 100644 --- a/packages/application/workspace/remove_member_use_case.py +++ b/packages/application/workspace/remove_member_use_case.py @@ -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)}" diff --git a/packages/application/workspace/subscription_use_case.py b/packages/application/workspace/subscription_use_case.py index c39ba6966..dd2952edd 100644 --- a/packages/application/workspace/subscription_use_case.py +++ b/packages/application/workspace/subscription_use_case.py @@ -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)}" diff --git a/packages/application/workspace/update_member_role_use_case.py b/packages/application/workspace/update_member_role_use_case.py index f7110b658..f11b0f208 100644 --- a/packages/application/workspace/update_member_role_use_case.py +++ b/packages/application/workspace/update_member_role_use_case.py @@ -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)}" diff --git a/packages/domain/__init__.py b/packages/domain/__init__.py index e8a1613a3..0a8dd79fa 100644 --- a/packages/domain/__init__.py +++ b/packages/domain/__init__.py @@ -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__ = [ diff --git a/packages/domain/auth/__init__.py b/packages/domain/auth/__init__.py index 12129fe3d..d3c9c52f3 100644 --- a/packages/domain/auth/__init__.py +++ b/packages/domain/auth/__init__.py @@ -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", diff --git a/packages/domain/auth/jwt_service.py b/packages/domain/auth/jwt_service.py index 8897391f1..0c8c59e53 100644 --- a/packages/domain/auth/jwt_service.py +++ b/packages/domain/auth/jwt_service.py @@ -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 diff --git a/packages/domain/auth/password_hasher.py b/packages/domain/auth/password_hasher.py index db468d9d8..6f7d3d155 100644 --- a/packages/domain/auth/password_hasher.py +++ b/packages/domain/auth/password_hasher.py @@ -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 diff --git a/packages/domain/classification.py b/packages/domain/classification.py index 07697de03..b9a63651c 100644 --- a/packages/domain/classification.py +++ b/packages/domain/classification.py @@ -28,6 +28,7 @@ class ClassificationJobStatus(StrEnum): class AssetClassification(StrEnum): """Asset classification categories.""" + SCENIC = "scenic" # 风景 PRODUCT = "product" # 产品 PERSON = "person" # 人物 diff --git a/packages/domain/entities.py b/packages/domain/entities.py index b56ff8d47..7d3ed3088 100644 --- a/packages/domain/entities.py +++ b/packages/domain/entities.py @@ -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 # 邀请人 diff --git a/packages/domain/permissions.py b/packages/domain/permissions.py index 096a2136c..0d2768d3b 100644 --- a/packages/domain/permissions.py +++ b/packages/domain/permissions.py @@ -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: 是否有权限 """ diff --git a/packages/domain/project_management.py b/packages/domain/project_management.py index 331bb3ccf..7a245b58b 100644 --- a/packages/domain/project_management.py +++ b/packages/domain/project_management.py @@ -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(), diff --git a/packages/domain/quota.py b/packages/domain/quota.py index aa20811f8..8db63e279 100644 --- a/packages/domain/quota.py +++ b/packages/domain/quota.py @@ -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: 警告级别 """ diff --git a/packages/ports/__init__.py b/packages/ports/__init__.py index f53fe2cb4..ece885bb6 100644 --- a/packages/ports/__init__.py +++ b/packages/ports/__init__.py @@ -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__ = [ diff --git a/packages/ports/generated_video_repository.py b/packages/ports/generated_video_repository.py index f062f878d..400107c2b 100644 --- a/packages/ports/generated_video_repository.py +++ b/packages/ports/generated_video_repository.py @@ -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]: ... diff --git a/packages/ports/generation_task_repository.py b/packages/ports/generation_task_repository.py index 84cc925a4..ffeac01dd 100644 --- a/packages/ports/generation_task_repository.py +++ b/packages/ports/generation_task_repository.py @@ -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: ... diff --git a/packages/ports/project_management_repositories.py b/packages/ports/project_management_repositories.py index 83e087faa..66dfa784b 100644 --- a/packages/ports/project_management_repositories.py +++ b/packages/ports/project_management_repositories.py @@ -1,4 +1,5 @@ """项目管理 Repository 接口定义""" + from abc import ABC, abstractmethod from packages.domain import Milestone, Task, TaskIssue diff --git a/packages/ports/project_repository.py b/packages/ports/project_repository.py index fc60b8452..82867e4e8 100644 --- a/packages/ports/project_repository.py +++ b/packages/ports/project_repository.py @@ -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: """删除项目""" diff --git a/packages/ports/user_repository.py b/packages/ports/user_repository.py index 68e9f4e6d..974ffd176 100644 --- a/packages/ports/user_repository.py +++ b/packages/ports/user_repository.py @@ -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: """删除用户""" diff --git a/packages/ports/workspace_invitation_repository.py b/packages/ports/workspace_invitation_repository.py index 8483bb937..78076e68a 100644 --- a/packages/ports/workspace_invitation_repository.py +++ b/packages/ports/workspace_invitation_repository.py @@ -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: """删除邀请""" diff --git a/packages/ports/workspace_member_repository.py b/packages/ports/workspace_member_repository.py index 4bf3ed436..ca71ff5c8 100644 --- a/packages/ports/workspace_member_repository.py +++ b/packages/ports/workspace_member_repository.py @@ -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: """删除成员""" diff --git a/packages/ports/workspace_repository.py b/packages/ports/workspace_repository.py index 6ee6529b0..d1a9c77fc 100644 --- a/packages/ports/workspace_repository.py +++ b/packages/ports/workspace_repository.py @@ -1,24 +1,26 @@ """ Workspace 仓储接口 """ + from abc import ABC, abstractmethod from typing import Optional + from packages.domain.entities import Workspace class WorkspaceRepository(ABC): """Workspace 仓储接口""" - + @abstractmethod def save(self, workspace: Workspace) -> None: """保存 Workspace""" pass - + @abstractmethod def find_by_id(self, workspace_id: str) -> Optional[Workspace]: """根据 ID 查找 Workspace""" pass - + @abstractmethod def delete(self, workspace_id: str) -> bool: """删除 Workspace""" diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 000000000..2ef261c59 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,7 @@ +[tool.black] +line-length = 120 +target-version = ["py312"] + +[tool.isort] +profile = "black" +line_length = 120 diff --git a/tests/integration/test_api.py b/tests/integration/test_api.py index 99a0aaf84..fae3059db 100644 --- a/tests/integration/test_api.py +++ b/tests/integration/test_api.py @@ -1,6 +1,7 @@ """ API 集成测试 """ + import pytest from fastapi.testclient import TestClient @@ -11,130 +12,165 @@ client = TestClient(app) class TestAuthAPI: """认证 API 集成测试""" - + def test_register_success(self): """测试注册成功""" - response = client.post("/api/v1/auth/register", json={ - "email": "test@example.com", - "password": "SecurePass123", - "username": "testuser", - "display_name": "Test User", - }) - + response = client.post( + "/api/v1/auth/register", + json={ + "email": "test@example.com", + "password": "SecurePass123", + "username": "testuser", + "display_name": "Test User", + }, + ) + assert response.status_code == 201 data = response.json() assert data["email"] == "test@example.com" assert data["username"] == "testuser" assert "user_id" in data - + def test_register_duplicate_email(self): """测试重复邮箱注册""" # 先注册一个用户 - client.post("/api/v1/auth/register", json={ - "email": "duplicate@example.com", - "password": "SecurePass123", - "username": "user1", - "display_name": "User 1", - }) - + client.post( + "/api/v1/auth/register", + json={ + "email": "duplicate@example.com", + "password": "SecurePass123", + "username": "user1", + "display_name": "User 1", + }, + ) + # 尝试用相同邮箱再次注册 - response = client.post("/api/v1/auth/register", json={ - "email": "duplicate@example.com", - "password": "SecurePass123", - "username": "user2", - "display_name": "User 2", - }) - + response = client.post( + "/api/v1/auth/register", + json={ + "email": "duplicate@example.com", + "password": "SecurePass123", + "username": "user2", + "display_name": "User 2", + }, + ) + assert response.status_code == 400 assert "already registered" in response.json()["detail"].lower() - + def test_login_success(self): """测试登录成功""" # 先注册 - client.post("/api/v1/auth/register", json={ - "email": "login@example.com", - "password": "SecurePass123", - "username": "loginuser", - "display_name": "Login User", - }) - + client.post( + "/api/v1/auth/register", + json={ + "email": "login@example.com", + "password": "SecurePass123", + "username": "loginuser", + "display_name": "Login User", + }, + ) + # 登录 - response = client.post("/api/v1/auth/login", json={ - "email": "login@example.com", - "password": "SecurePass123", - }) - + response = client.post( + "/api/v1/auth/login", + json={ + "email": "login@example.com", + "password": "SecurePass123", + }, + ) + assert response.status_code == 200 data = response.json() assert "access_token" in data assert "refresh_token" in data assert data["token_type"] == "bearer" - + def test_login_wrong_password(self): """测试密码错误""" - response = client.post("/api/v1/auth/login", json={ - "email": "login@example.com", - "password": "WrongPassword123", - }) - + response = client.post( + "/api/v1/auth/login", + json={ + "email": "login@example.com", + "password": "WrongPassword123", + }, + ) + assert response.status_code == 401 class TestWorkspaceAPI: """工作空间 API 集成测试""" - + def setup_method(self): """每个测试前的准备""" # 注册并登录,获取 token - client.post("/api/v1/auth/register", json={ - "email": "workspace@example.com", - "password": "SecurePass123", - "username": "workspaceuser", - "display_name": "Workspace User", - }) - - response = client.post("/api/v1/auth/login", json={ - "email": "workspace@example.com", - "password": "SecurePass123", - }) - + client.post( + "/api/v1/auth/register", + json={ + "email": "workspace@example.com", + "password": "SecurePass123", + "username": "workspaceuser", + "display_name": "Workspace User", + }, + ) + + response = client.post( + "/api/v1/auth/login", + json={ + "email": "workspace@example.com", + "password": "SecurePass123", + }, + ) + self.token = response.json()["access_token"] self.headers = {"Authorization": f"Bearer {self.token}"} - + def test_create_workspace(self): """测试创建工作空间""" - response = client.post("/api/v1/workspaces", json={ - "name": "My Workspace", - "subscription_plan": "free", - }, headers=self.headers) - + response = client.post( + "/api/v1/workspaces", + json={ + "name": "My Workspace", + "subscription_plan": "free", + }, + headers=self.headers, + ) + assert response.status_code == 201 data = response.json() assert data["name"] == "My Workspace" assert data["subscription_plan"] == "free" assert data["max_projects"] == 3 - + def test_list_workspaces(self): """测试获取工作空间列表""" # 创建工作空间 - client.post("/api/v1/workspaces", json={ - "name": "Workspace 1", - }, headers=self.headers) - + client.post( + "/api/v1/workspaces", + json={ + "name": "Workspace 1", + }, + headers=self.headers, + ) + # 获取列表 response = client.get("/api/v1/workspaces", headers=self.headers) - + assert response.status_code == 200 data = response.json() assert len(data["workspaces"]) > 0 assert data["workspaces"][0]["name"] == "Workspace 1" - + def test_create_workspace_unauthorized(self): """测试未登录创建工作空间""" - response = client.post("/api/v1/workspaces", json={ - "name": "Unauthorized Workspace", - }) - + response = client.post( + "/api/v1/workspaces", + json={ + "name": "Unauthorized Workspace", + }, + ) + assert response.status_code == 403 # FastAPI HTTPBearer 返回 403 diff --git a/tests/integration/test_asset_tags.py b/tests/integration/test_asset_tags.py index 5fc4c5871..a45cb24f8 100644 --- a/tests/integration/test_asset_tags.py +++ b/tests/integration/test_asset_tags.py @@ -13,10 +13,10 @@ def test_add_tag_to_asset(): storage_key="uploads/abc/video.mp4", mime_type="video/mp4", ) - + asset.add_tag("风景") asset.add_tag("自然") - + assert len(asset.tags) == 2 assert "风景" in asset.tags assert "自然" in asset.tags @@ -32,10 +32,10 @@ def test_add_duplicate_tag_should_ignore(): storage_key="uploads/abc/video.mp4", mime_type="video/mp4", ) - + asset.add_tag("风景") asset.add_tag("风景") # 重复 - + assert len(asset.tags) == 1 assert asset.tags.count("风景") == 1 @@ -50,10 +50,10 @@ def test_add_empty_tag_should_fail(): storage_key="uploads/abc/video.mp4", mime_type="video/mp4", ) - + with pytest.raises(ValueError, match="标签不能为空"): asset.add_tag("") - + with pytest.raises(ValueError, match="标签不能为空"): asset.add_tag(" ") # 仅空格 @@ -68,12 +68,12 @@ def test_remove_tag_from_asset(): storage_key="uploads/abc/video.mp4", mime_type="video/mp4", ) - + asset.add_tag("风景") asset.add_tag("自然") - + asset.remove_tag("风景") - + assert len(asset.tags) == 1 assert "风景" not in asset.tags assert "自然" in asset.tags @@ -89,11 +89,11 @@ def test_remove_nonexistent_tag_should_be_idempotent(): storage_key="uploads/abc/video.mp4", mime_type="video/mp4", ) - + asset.add_tag("风景") - + # 删除不存在的标签,不应报错 asset.remove_tag("不存在的标签") - + assert len(asset.tags) == 1 assert "风景" in asset.tags diff --git a/tests/integration/test_classification_pipeline.py b/tests/integration/test_classification_pipeline.py index 879fecc22..7564f3469 100644 --- a/tests/integration/test_classification_pipeline.py +++ b/tests/integration/test_classification_pipeline.py @@ -1,6 +1,9 @@ -from packages.application import SubmitClassificationJobCommand, SubmitClassificationJobUseCase from packages.adapters.in_memory import InMemoryClassificationJobRepository -from packages.domain import ClassificationJobStatus, AssetClassification +from packages.application import ( + SubmitClassificationJobCommand, + SubmitClassificationJobUseCase, +) +from packages.domain import AssetClassification, ClassificationJobStatus def simulate_classify_asset(job_id: str, job_repo: InMemoryClassificationJobRepository) -> dict: @@ -8,24 +11,24 @@ def simulate_classify_asset(job_id: str, job_repo: InMemoryClassificationJobRepo 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) - + # Mock classification 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) - + return { "status": "completed", "job_id": job.id, @@ -37,7 +40,7 @@ def simulate_classify_asset(job_id: str, job_repo: InMemoryClassificationJobRepo job.status = ClassificationJobStatus.FAILED job.error_message = str(e) job_repo.update(job) - + return { "status": "failed", "job_id": job.id, @@ -48,7 +51,7 @@ def simulate_classify_asset(job_id: str, job_repo: InMemoryClassificationJobRepo def test_classification_pipeline(): """Test the full classification pipeline: submit job -> worker processes -> result.""" job_repo = InMemoryClassificationJobRepository() - + # Submit classification job use_case = SubmitClassificationJobUseCase(job_repo) job = use_case.execute( @@ -58,18 +61,18 @@ def test_classification_pipeline(): asset_id="asset-123", ) ) - + assert job.status == ClassificationJobStatus.PENDING assert job.classification == "" assert job.confidence == 0.0 - + # Simulate worker task execution result = simulate_classify_asset(job.id, job_repo) - + assert result["status"] == "completed" assert "classification" in result assert "confidence" in result - + # Verify job was updated updated_job = job_repo.get(job.id) assert updated_job is not None diff --git a/tests/integration/test_generation_pipeline.py b/tests/integration/test_generation_pipeline.py index 0a71879a8..bf2c7c77b 100644 --- a/tests/integration/test_generation_pipeline.py +++ b/tests/integration/test_generation_pipeline.py @@ -1,6 +1,10 @@ from datetime import datetime, timezone -from packages.application import CreateGenerationTaskCommand, CreateGenerationTaskUseCase, GetGeneratedVideoDownloadUrlUseCase +from packages.application import ( + CreateGenerationTaskCommand, + CreateGenerationTaskUseCase, + GetGeneratedVideoDownloadUrlUseCase, +) from packages.domain import GeneratedVideo, GenerationTaskStatus @@ -41,7 +45,11 @@ class DummyGeneratedVideoRepository: return [video for video in self.items.values() if video.generation_task_id == generation_task_id] -def simulate_generate_video(task_id: str, task_repo: DummyGenerationTaskRepository, video_repo: DummyGeneratedVideoRepository) -> dict: +def simulate_generate_video( + task_id: str, + task_repo: DummyGenerationTaskRepository, + video_repo: DummyGeneratedVideoRepository, +) -> dict: task = task_repo.get(task_id) if task is None: return {"status": "failed", "error": "task not found"} @@ -75,7 +83,12 @@ def simulate_generate_video(task_id: str, task_repo: DummyGenerationTaskReposito 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, + } def test_create_generation_task_smoke(): diff --git a/tests/integration/test_ingest_pipeline.py b/tests/integration/test_ingest_pipeline.py index e6526a8fe..9a62ac1ce 100644 --- a/tests/integration/test_ingest_pipeline.py +++ b/tests/integration/test_ingest_pipeline.py @@ -1,9 +1,16 @@ +from packages.adapters.in_memory import ( + InMemoryAssetRepository, + InMemoryIngestJobRepository, +) from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase -from packages.adapters.in_memory import InMemoryAssetRepository, InMemoryIngestJobRepository from packages.domain import Asset, IngestJob, IngestJobStatus -def simulate_ingest_asset(job_id: str, job_repo: InMemoryIngestJobRepository, asset_repo: InMemoryAssetRepository) -> dict: +def simulate_ingest_asset( + job_id: str, + job_repo: InMemoryIngestJobRepository, + asset_repo: InMemoryAssetRepository, +) -> dict: """ Simulate ingest asset logic without Celery. This is the core business logic that would run inside the worker task. @@ -11,12 +18,12 @@ def simulate_ingest_asset(job_id: str, job_repo: InMemoryIngestJobRepository, as 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_repo.update(job) - + # Mock metadata extraction mime_type = "video/mp4" if job.storage_key.endswith(".mp4") else "image/jpeg" metadata = { @@ -25,10 +32,10 @@ def simulate_ingest_asset(job_id: str, job_repo: InMemoryIngestJobRepository, as "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, @@ -40,12 +47,12 @@ def simulate_ingest_asset(job_id: str, job_repo: InMemoryIngestJobRepository, as metadata=metadata, ) asset_repo.create(asset) - + # Update job status to COMPLETED job.status = IngestJobStatus.COMPLETED job.result_asset_id = asset.id job_repo.update(job) - + return { "status": "completed", "job_id": job.id, @@ -56,7 +63,7 @@ def simulate_ingest_asset(job_id: str, job_repo: InMemoryIngestJobRepository, as job.status = IngestJobStatus.FAILED job.error_message = str(e) job_repo.update(job) - + return { "status": "failed", "job_id": job.id, @@ -68,7 +75,7 @@ def test_ingest_asset_pipeline(): """Test the full ingest pipeline: submit job -> worker processes -> asset created.""" job_repo = InMemoryIngestJobRepository() asset_repo = InMemoryAssetRepository() - + # Submit ingest job use_case = SubmitIngestJobUseCase(job_repo) job = use_case.execute( @@ -79,22 +86,22 @@ def test_ingest_asset_pipeline(): storage_key="uploads/test-video.mp4", ) ) - + assert job.status == IngestJobStatus.PENDING assert job.result_asset_id == "" - + # Simulate worker task execution result = simulate_ingest_asset(job.id, job_repo, asset_repo) - + assert result["status"] == "completed" assert "asset_id" in result - + # Verify job was updated updated_job = job_repo.get(job.id) assert updated_job is not None assert updated_job.status == IngestJobStatus.COMPLETED assert updated_job.result_asset_id != "" - + # Verify asset was created assets = asset_repo.list_by_library("lib-1") assert len(assets) == 1 diff --git a/tests/integration/test_project_management.py b/tests/integration/test_project_management.py index b37003e66..d86606272 100644 --- a/tests/integration/test_project_management.py +++ b/tests/integration/test_project_management.py @@ -1,4 +1,5 @@ """项目管理功能集成测试""" + import pytest from packages.adapters.in_memory.project_management_repositories import ( @@ -23,7 +24,7 @@ def test_create_task(): """测试创建任务""" repo = InMemoryTaskRepository() use_case = CreateTaskUseCase(repo) - + task = use_case.execute( project_id="proj_1", workspace_id="ws_1", @@ -31,7 +32,7 @@ def test_create_task(): description="实现用户登录功能", priority=TaskPriority.HIGH, ) - + assert task.id is not None assert task.name == "开发登录功能" assert task.status == TaskStatus.PENDING @@ -43,7 +44,7 @@ def test_list_tasks(): """测试获取任务列表""" repo = InMemoryTaskRepository() create_use_case = CreateTaskUseCase(repo) - + # 创建两个任务 create_use_case.execute( project_id="proj_1", @@ -55,11 +56,11 @@ def test_list_tasks(): workspace_id="ws_1", name="任务2", ) - + # 查询任务列表 list_use_case = ListProjectTasksUseCase(repo) tasks = list_use_case.execute("proj_1") - + assert len(tasks) == 2 assert tasks[0].name == "任务1" assert tasks[1].name == "任务2" @@ -70,17 +71,17 @@ def test_update_task_status(): repo = InMemoryTaskRepository() create_use_case = CreateTaskUseCase(repo) update_use_case = UpdateTaskStatusUseCase(repo) - + # 创建任务 task = create_use_case.execute( project_id="proj_1", workspace_id="ws_1", name="测试任务", ) - + # 更新状态为进行中 updated_task = update_use_case.execute(task.id, TaskStatus.IN_PROGRESS) - + assert updated_task.status == TaskStatus.IN_PROGRESS assert updated_task.actual_start_date is not None @@ -90,23 +91,23 @@ def test_update_task_progress(): repo = InMemoryTaskRepository() create_use_case = CreateTaskUseCase(repo) progress_use_case = UpdateTaskProgressUseCase(repo) - + # 创建任务 task = create_use_case.execute( project_id="proj_1", workspace_id="ws_1", name="测试任务", ) - + # 更新进度到 50% updated_task = progress_use_case.execute(task.id, 50.0) - + assert updated_task.progress == 50.0 assert updated_task.status == TaskStatus.IN_PROGRESS - + # 更新进度到 100% completed_task = progress_use_case.execute(task.id, 100.0) - + assert completed_task.progress == 100.0 assert completed_task.status == TaskStatus.COMPLETED assert completed_task.actual_end_date is not None @@ -116,14 +117,14 @@ def test_create_milestone(): """测试创建里程碑""" repo = InMemoryMilestoneRepository() use_case = CreateMilestoneUseCase(repo) - + milestone = use_case.execute( project_id="proj_1", workspace_id="ws_1", name="V1.0 发布", description="第一个正式版本", ) - + assert milestone.id is not None assert milestone.name == "V1.0 发布" assert milestone.completed is False @@ -135,7 +136,7 @@ def test_create_and_resolve_issue(): create_use_case = CreateTaskIssueUseCase(repo) resolve_use_case = ResolveTaskIssueUseCase(repo) list_use_case = ListTaskIssuesUseCase(repo) - + # 创建问题 issue = create_use_case.execute( task_id="task_1", @@ -144,17 +145,17 @@ def test_create_and_resolve_issue(): title="接口报错", description="调用登录接口返回 500", ) - + assert issue.id is not None assert issue.title == "接口报错" assert issue.resolved is False - + # 解决问题 resolved_issue = resolve_use_case.execute(issue.id) - + assert resolved_issue.resolved is True assert resolved_issue.resolved_at is not None - + # 查询任务问题列表 issues = list_use_case.execute("task_1") assert len(issues) == 1 @@ -165,14 +166,14 @@ def test_task_hierarchy(): """测试任务层级关系""" repo = InMemoryTaskRepository() create_use_case = CreateTaskUseCase(repo) - + # 创建父任务 parent_task = create_use_case.execute( project_id="proj_1", workspace_id="ws_1", name="开发用户模块", ) - + # 创建子任务 child_task_1 = create_use_case.execute( project_id="proj_1", @@ -180,17 +181,17 @@ def test_task_hierarchy(): name="登录功能", parent_task_id=parent_task.id, ) - + child_task_2 = create_use_case.execute( project_id="proj_1", workspace_id="ws_1", name="注册功能", parent_task_id=parent_task.id, ) - + # 查询子任务 children = repo.list_by_parent(parent_task.id) - + assert len(children) == 2 assert children[0].parent_task_id == parent_task.id assert children[1].parent_task_id == parent_task.id @@ -199,11 +200,11 @@ def test_task_hierarchy(): def test_get_task_detail(): """测试获取任务详情""" from packages.application.get_task_detail_use_case import GetTaskDetailUseCase - + repo = InMemoryTaskRepository() create_use_case = CreateTaskUseCase(repo) get_use_case = GetTaskDetailUseCase(repo) - + # 创建任务 task = create_use_case.execute( project_id="proj_1", @@ -211,14 +212,14 @@ def test_get_task_detail(): name="测试任务", description="这是一个测试任务", ) - + # 获取详情 retrieved_task = get_use_case.execute(task.id) - + assert retrieved_task.id == task.id assert retrieved_task.name == "测试任务" assert retrieved_task.description == "这是一个测试任务" - + # 测试不存在的任务 try: get_use_case.execute("nonexistent_id") @@ -230,11 +231,11 @@ def test_get_task_detail(): def test_update_task(): """测试任务基本信息更新""" from packages.application.update_task_use_case import UpdateTaskUseCase - + repo = InMemoryTaskRepository() create_use_case = CreateTaskUseCase(repo) update_use_case = UpdateTaskUseCase(repo) - + # 创建任务 task = create_use_case.execute( project_id="proj_1", @@ -243,7 +244,7 @@ def test_update_task(): description="原始描述", priority="low", ) - + # 更新任务 updated_task = update_use_case.execute( task_id=task.id, @@ -251,17 +252,17 @@ def test_update_task(): description="更新后的描述", priority="high", ) - + assert updated_task.name == "更新后的任务" assert updated_task.description == "更新后的描述" assert updated_task.priority == "high" - + # 部分更新 partial_updated = update_use_case.execute( task_id=task.id, name="又更新了", ) - + assert partial_updated.name == "又更新了" assert partial_updated.description == "更新后的描述" # 保持不变 assert partial_updated.priority == "high" # 保持不变 diff --git a/tests/integration/test_projects.py b/tests/integration/test_projects.py index 7ca5e5862..77342dfcf 100644 --- a/tests/integration/test_projects.py +++ b/tests/integration/test_projects.py @@ -1,3 +1,9 @@ +from packages.adapters.in_memory import ( + InMemoryAssetLibraryRepository, + InMemoryAssetRepository, + InMemoryIngestJobRepository, + InMemoryProjectRepository, +) from packages.application import ( CreateAssetCommand, CreateAssetLibraryCommand, @@ -11,12 +17,6 @@ from packages.application import ( SubmitIngestJobCommand, SubmitIngestJobUseCase, ) -from packages.adapters.in_memory import ( - InMemoryAssetLibraryRepository, - InMemoryAssetRepository, - InMemoryIngestJobRepository, - InMemoryProjectRepository, -) from packages.domain import AssetLibraryKind, IngestJobStatus diff --git a/tests/integration/test_sqlalchemy_repositories.py b/tests/integration/test_sqlalchemy_repositories.py index c030a4caa..4c01d9fa0 100644 --- a/tests/integration/test_sqlalchemy_repositories.py +++ b/tests/integration/test_sqlalchemy_repositories.py @@ -2,7 +2,9 @@ from sqlalchemy import create_engine from sqlalchemy.orm import Session, sessionmaker from packages.adapters.sqlalchemy_impl.models import Base -from packages.adapters.sqlalchemy_impl.project_repository import SQLAlchemyProjectRepository +from packages.adapters.sqlalchemy_impl.project_repository import ( + SQLAlchemyProjectRepository, +) from packages.application import CreateProjectCommand, CreateProjectUseCase @@ -13,12 +15,12 @@ def test_sqlalchemy_project_repository(): Base.metadata.create_all(engine) SessionLocal = sessionmaker(bind=engine) session: Session = SessionLocal() - + try: # Create repository and use case repository = SQLAlchemyProjectRepository(session) use_case = CreateProjectUseCase(repository) - + # Create project project = use_case.execute( CreateProjectCommand( @@ -27,10 +29,10 @@ def test_sqlalchemy_project_repository(): description="Test description", ) ) - + assert project.name == "Test Project" assert project.workspace_id == "ws-1" - + # List projects projects = repository.list_by_workspace("ws-1") assert len(projects) == 1 diff --git a/tests/integration/test_upload_pipeline.py b/tests/integration/test_upload_pipeline.py index a70a99015..9cee070fa 100644 --- a/tests/integration/test_upload_pipeline.py +++ b/tests/integration/test_upload_pipeline.py @@ -1,5 +1,8 @@ +from packages.adapters.in_memory import ( + InMemoryAssetRepository, + InMemoryIngestJobRepository, +) from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase -from packages.adapters.in_memory import InMemoryAssetRepository, InMemoryIngestJobRepository from packages.domain import IngestJobStatus @@ -13,11 +16,12 @@ def simulate_upload_and_ingest( ) -> dict: """Simulate full upload → ingest pipeline.""" from uuid import uuid4 + from tests.integration.test_ingest_pipeline import simulate_ingest_asset - + # Mock storage: generate storage_key storage_key = f"uploads/{uuid4().hex[:8]}/{filename}" - + # Submit ingest job use_case = SubmitIngestJobUseCase(job_repo) job = use_case.execute( @@ -28,10 +32,10 @@ def simulate_upload_and_ingest( storage_key=storage_key, ) ) - + # Simulate worker task result = simulate_ingest_asset(job.id, job_repo, asset_repo) - + return { "storage_key": storage_key, "job_id": job.id, @@ -43,7 +47,7 @@ def test_upload_to_asset_full_pipeline(): """Test full pipeline: upload → storage → ingest job → worker → asset created.""" job_repo = InMemoryIngestJobRepository() asset_repo = InMemoryAssetRepository() - + # Simulate upload result = simulate_upload_and_ingest( workspace_id="ws-1", @@ -53,17 +57,17 @@ def test_upload_to_asset_full_pipeline(): job_repo=job_repo, asset_repo=asset_repo, ) - + assert "storage_key" in result assert result["storage_key"].endswith("demo-video.mp4") assert result["worker_result"]["status"] == "completed" - + # Verify job was completed job = job_repo.get(result["job_id"]) assert job is not None assert job.status == IngestJobStatus.COMPLETED assert job.result_asset_id != "" - + # Verify asset was created assets = asset_repo.list_by_library("lib-1") assert len(assets) == 1 diff --git a/tests/unit/test_accept_invitation_use_case.py b/tests/unit/test_accept_invitation_use_case.py index 43f7b113f..3c7bcc8f4 100644 --- a/tests/unit/test_accept_invitation_use_case.py +++ b/tests/unit/test_accept_invitation_use_case.py @@ -1,62 +1,71 @@ """ 接受/拒绝邀请 Use Case 测试 """ -import pytest -from unittest.mock import Mock + from datetime import datetime, timedelta, timezone +from unittest.mock import Mock + +import pytest + from packages.application.workspace.accept_invitation_use_case import ( - AcceptInvitationUseCase, AcceptInvitationRequest, - DeclineInvitationUseCase, + AcceptInvitationUseCase, DeclineInvitationRequest, + DeclineInvitationUseCase, ) from packages.domain.entities import ( + InvitationStatus, + User, Workspace, WorkspaceInvitation, WorkspaceMember, - User, - InvitationStatus, ) class TestAcceptInvitationUseCase: """接受邀请测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_workspace_and_user = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def mock_invitation_repo(self): repo = Mock() repo.find_by_token = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def mock_user_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture - def use_case(self, mock_workspace_repo, mock_member_repo, mock_invitation_repo, mock_user_repo): + def use_case( + self, + mock_workspace_repo, + mock_member_repo, + mock_invitation_repo, + mock_user_repo, + ): return AcceptInvitationUseCase( workspace_repository=mock_workspace_repo, workspace_member_repository=mock_member_repo, workspace_invitation_repository=mock_invitation_repo, user_repository=mock_user_repo, ) - + @pytest.fixture def test_workspace(self): return Workspace( @@ -64,7 +73,7 @@ class TestAcceptInvitationUseCase: name="Test Workspace", owner_user_id="owner-id", ) - + @pytest.fixture def test_user(self): return User( @@ -73,7 +82,7 @@ class TestAcceptInvitationUseCase: username="invitee", display_name="Invitee User", ) - + @pytest.fixture def valid_invitation(self): return WorkspaceInvitation( @@ -86,7 +95,7 @@ class TestAcceptInvitationUseCase: status=InvitationStatus.PENDING, expires_at=datetime.now(timezone.utc) + timedelta(days=7), ) - + def test_accept_invitation_success( self, use_case, @@ -103,45 +112,45 @@ class TestAcceptInvitationUseCase: mock_user_repo.find_by_id.return_value = test_user mock_workspace_repo.find_by_id.return_value = test_workspace mock_member_repo.find_by_workspace_and_user.return_value = None - + request = AcceptInvitationRequest( invitation_token="valid-token", user_id="user-123", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.workspace_id == "workspace-123" assert response.workspace_name == "Test Workspace" assert response.role == "member" - + # 验证创建了成员记录 mock_member_repo.save.assert_called_once() member = mock_member_repo.save.call_args[0][0] assert member.user_id == "user-123" assert member.role == "member" assert member.invited_by == "inviter-id" - + # 验证更新了邀请状态 assert valid_invitation.status == InvitationStatus.ACCEPTED assert valid_invitation.accepted_at is not None - + def test_accept_invitation_invalid_token(self, use_case, mock_invitation_repo): """测试无效令牌""" mock_invitation_repo.find_by_token.return_value = None - + request = AcceptInvitationRequest( invitation_token="invalid-token", user_id="user-123", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Invalid invitation token" - + def test_accept_invitation_already_accepted( self, use_case, @@ -151,17 +160,17 @@ class TestAcceptInvitationUseCase: """测试邀请已被接受""" valid_invitation.status = InvitationStatus.ACCEPTED mock_invitation_repo.find_by_token.return_value = valid_invitation - + request = AcceptInvitationRequest( invitation_token="valid-token", user_id="user-123", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Invitation has already been accepted" - + def test_accept_invitation_expired( self, use_case, @@ -171,18 +180,18 @@ class TestAcceptInvitationUseCase: """测试邀请已过期""" valid_invitation.expires_at = datetime.now(timezone.utc) - timedelta(days=1) mock_invitation_repo.find_by_token.return_value = valid_invitation - + request = AcceptInvitationRequest( invitation_token="valid-token", user_id="user-123", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Invitation has expired" assert valid_invitation.status == InvitationStatus.EXPIRED - + def test_accept_invitation_email_mismatch( self, use_case, @@ -192,7 +201,7 @@ class TestAcceptInvitationUseCase: ): """测试邮箱不匹配""" mock_invitation_repo.find_by_token.return_value = valid_invitation - + different_user = User( id="user-123", email="different@test.com", @@ -200,17 +209,17 @@ class TestAcceptInvitationUseCase: display_name="Different User", ) mock_user_repo.find_by_id.return_value = different_user - + request = AcceptInvitationRequest( invitation_token="valid-token", user_id="user-123", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "This invitation is for a different email address" - + def test_accept_invitation_already_member( self, use_case, @@ -226,7 +235,7 @@ class TestAcceptInvitationUseCase: mock_invitation_repo.find_by_token.return_value = valid_invitation mock_user_repo.find_by_id.return_value = test_user mock_workspace_repo.find_by_id.return_value = test_workspace - + existing_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -234,41 +243,41 @@ class TestAcceptInvitationUseCase: role="admin", ) mock_member_repo.find_by_workspace_and_user.return_value = existing_member - + request = AcceptInvitationRequest( invitation_token="valid-token", user_id="user-123", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.role == "admin" # 返回现有角色 - + # 不创建新成员记录 mock_member_repo.save.assert_not_called() - + # 但仍标记邀请为已接受 assert valid_invitation.status == InvitationStatus.ACCEPTED class TestDeclineInvitationUseCase: """拒绝邀请测试""" - + @pytest.fixture def mock_invitation_repo(self): repo = Mock() repo.find_by_token = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def use_case(self, mock_invitation_repo): return DeclineInvitationUseCase( workspace_invitation_repository=mock_invitation_repo, ) - + @pytest.fixture def valid_invitation(self): return WorkspaceInvitation( @@ -281,35 +290,35 @@ class TestDeclineInvitationUseCase: status=InvitationStatus.PENDING, expires_at=datetime.now(timezone.utc) + timedelta(days=7), ) - + def test_decline_invitation_success(self, use_case, mock_invitation_repo, valid_invitation): """测试拒绝邀请成功""" mock_invitation_repo.find_by_token.return_value = valid_invitation - + request = DeclineInvitationRequest(invitation_token="valid-token") success, error = use_case.execute(request) - + assert success is True assert error is None assert valid_invitation.status == InvitationStatus.DECLINED - + def test_decline_invitation_invalid_token(self, use_case, mock_invitation_repo): """测试无效令牌""" mock_invitation_repo.find_by_token.return_value = None - + request = DeclineInvitationRequest(invitation_token="invalid-token") success, error = use_case.execute(request) - + assert success is False assert error == "Invalid invitation token" - + def test_decline_invitation_already_accepted(self, use_case, mock_invitation_repo, valid_invitation): """测试邀请已被接受""" valid_invitation.status = InvitationStatus.ACCEPTED mock_invitation_repo.find_by_token.return_value = valid_invitation - + request = DeclineInvitationRequest(invitation_token="valid-token") success, error = use_case.execute(request) - + assert success is False assert error == "Invitation has already been accepted" diff --git a/tests/unit/test_auth_simple.py b/tests/unit/test_auth_simple.py index d4541ff41..30a5e9b0c 100644 --- a/tests/unit/test_auth_simple.py +++ b/tests/unit/test_auth_simple.py @@ -18,6 +18,7 @@ spec.loader.exec_module(auth_simple) _create_access_token = auth_simple._create_access_token _verify_password_with_legacy_upgrade = auth_simple._verify_password_with_legacy_upgrade from app.config import settings + from packages.adapters.sqlalchemy_impl.models import UserModel from packages.domain.auth import password_hasher diff --git a/tests/unit/test_create_workspace_use_case.py b/tests/unit/test_create_workspace_use_case.py index e015b3afc..f72bcf50c 100644 --- a/tests/unit/test_create_workspace_use_case.py +++ b/tests/unit/test_create_workspace_use_case.py @@ -1,36 +1,39 @@ """ 创建 Workspace Use Case 测试 """ -import pytest + from unittest.mock import Mock + +import pytest + from packages.application.workspace import ( - CreateWorkspaceUseCase, CreateWorkspaceRequest, + CreateWorkspaceUseCase, ) from packages.domain.entities import User class TestCreateWorkspaceUseCase: """创建工作空间测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.save = Mock() return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.save = Mock() return repo - + @pytest.fixture def mock_user_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture def use_case(self, mock_workspace_repo, mock_member_repo, mock_user_repo): return CreateWorkspaceUseCase( @@ -38,7 +41,7 @@ class TestCreateWorkspaceUseCase: workspace_member_repository=mock_member_repo, user_repository=mock_user_repo, ) - + @pytest.fixture def test_user(self): return User( @@ -47,123 +50,129 @@ class TestCreateWorkspaceUseCase: username="testuser", display_name="Test User", ) - - def test_create_workspace_success_free_plan(self, use_case, mock_workspace_repo, mock_member_repo, mock_user_repo, test_user): + + def test_create_workspace_success_free_plan( + self, use_case, mock_workspace_repo, mock_member_repo, mock_user_repo, test_user + ): """测试创建免费工作空间""" mock_user_repo.find_by_id.return_value = test_user - + request = CreateWorkspaceRequest( name="My Workspace", owner_user_id="user-123", subscription_plan="free", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.name == "My Workspace" assert response.subscription_plan == "free" assert response.max_projects == 3 assert response.max_storage_gb == 10 - + # 验证保存了 workspace mock_workspace_repo.save.assert_called_once() workspace = mock_workspace_repo.save.call_args[0][0] assert workspace.name == "My Workspace" assert workspace.owner_user_id == "user-123" - + # 验证创建了 owner 成员 mock_member_repo.save.assert_called_once() member = mock_member_repo.save.call_args[0][0] assert member.user_id == "user-123" assert member.role == "owner" - - def test_create_workspace_pro_plan(self, use_case, mock_workspace_repo, mock_member_repo, mock_user_repo, test_user): + + def test_create_workspace_pro_plan( + self, use_case, mock_workspace_repo, mock_member_repo, mock_user_repo, test_user + ): """测试创建 Pro 工作空间""" mock_user_repo.find_by_id.return_value = test_user - + request = CreateWorkspaceRequest( name="Pro Workspace", owner_user_id="user-123", subscription_plan="pro", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.subscription_plan == "pro" assert response.max_projects == 999999 # unlimited assert response.max_storage_gb == 100 - - def test_create_workspace_enterprise_plan(self, use_case, mock_workspace_repo, mock_member_repo, mock_user_repo, test_user): + + def test_create_workspace_enterprise_plan( + self, use_case, mock_workspace_repo, mock_member_repo, mock_user_repo, test_user + ): """测试创建 Enterprise 工作空间""" mock_user_repo.find_by_id.return_value = test_user - + request = CreateWorkspaceRequest( name="Enterprise Workspace", owner_user_id="user-123", subscription_plan="enterprise", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.subscription_plan == "enterprise" assert response.max_projects == 999999 assert response.max_storage_gb == 1000 - + def test_create_workspace_missing_name(self, use_case): """测试缺少名称""" request = CreateWorkspaceRequest( name="", owner_user_id="user-123", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Workspace name is required" - + def test_create_workspace_name_too_long(self, use_case): """测试名称过长""" request = CreateWorkspaceRequest( name="A" * 101, owner_user_id="user-123", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Workspace name is too long (max 100 characters)" - + def test_create_workspace_user_not_found(self, use_case, mock_user_repo): """测试用户不存在""" mock_user_repo.find_by_id.return_value = None - + request = CreateWorkspaceRequest( name="My Workspace", owner_user_id="nonexistent-user", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Owner user not found" - + def test_create_workspace_invalid_plan(self, use_case, mock_user_repo, test_user): """测试无效的订阅计划""" mock_user_repo.find_by_id.return_value = test_user - + request = CreateWorkspaceRequest( name="My Workspace", owner_user_id="user-123", subscription_plan="invalid_plan", ) - + response, error = use_case.execute(request) - + assert response is None assert "Invalid subscription plan" in error diff --git a/tests/unit/test_email_service.py b/tests/unit/test_email_service.py index 116687fef..2bbbe74dc 100644 --- a/tests/unit/test_email_service.py +++ b/tests/unit/test_email_service.py @@ -1,14 +1,17 @@ """ 邮件服务测试 """ + +from unittest.mock import MagicMock, Mock, patch + import pytest -from unittest.mock import Mock, patch, MagicMock -from packages.domain.auth.email_service import EmailService, EmailConfig + +from packages.domain.auth.email_service import EmailConfig, EmailService class TestEmailService: """邮件服务测试""" - + @pytest.fixture def email_config(self): """创建测试邮件配置""" @@ -21,41 +24,41 @@ class TestEmailService: from_name="Test Service", use_tls=True, ) - + @pytest.fixture def email_service(self, email_config): """创建邮件服务实例""" return EmailService(config=email_config) - - @patch('smtplib.SMTP') + + @patch("smtplib.SMTP") def test_send_email_success(self, mock_smtp, email_service): """测试发送邮件成功""" # Mock SMTP 服务器 mock_server = MagicMock() mock_smtp.return_value.__enter__.return_value = mock_server - + success, error = email_service.send_email( to_email="user@example.com", subject="Test Email", html_body="

Test

", text_body="Test", ) - + assert success is True assert error is None - + # 验证 SMTP 调用 mock_smtp.assert_called_once_with("smtp.test.com", 587) mock_server.starttls.assert_called_once() mock_server.login.assert_called_once_with("test@test.com", "test-password") mock_server.sendmail.assert_called_once() - - @patch('smtplib.SMTP') + + @patch("smtplib.SMTP") def test_send_email_with_cc_bcc(self, mock_smtp, email_service): """测试发送邮件带抄送和密送""" mock_server = MagicMock() mock_smtp.return_value.__enter__.return_value = mock_server - + success, error = email_service.send_email( to_email="user@example.com", subject="Test Email", @@ -63,9 +66,9 @@ class TestEmailService: cc=["cc1@example.com", "cc2@example.com"], bcc=["bcc@example.com"], ) - + assert success is True - + # 验证收件人列表包含所有人 call_args = mock_server.sendmail.call_args recipients = call_args[0][1] @@ -73,73 +76,73 @@ class TestEmailService: assert "cc1@example.com" in recipients assert "cc2@example.com" in recipients assert "bcc@example.com" in recipients - - @patch('smtplib.SMTP') + + @patch("smtplib.SMTP") def test_send_email_smtp_error(self, mock_smtp, email_service): """测试 SMTP 错误处理""" # Mock SMTP 抛出异常 mock_smtp.side_effect = Exception("SMTP connection failed") - + success, error = email_service.send_email( to_email="user@example.com", subject="Test", html_body="

Test

", ) - + assert success is False assert error is not None assert "SMTP connection failed" in error - - @patch('smtplib.SMTP') + + @patch("smtplib.SMTP") def test_send_verification_email(self, mock_smtp, email_service): """测试发送验证邮件""" mock_server = MagicMock() mock_smtp.return_value.__enter__.return_value = mock_server - + success, error = email_service.send_verification_email( to_email="user@example.com", username="TestUser", verification_url="https://example.com/verify?token=abc123", ) - + assert success is True assert error is None - + # 验证发送了邮件 mock_server.sendmail.assert_called_once() call_args = mock_server.sendmail.call_args - + # 验证收件人 assert call_args[0][1] == ["user@example.com"] - - @patch('smtplib.SMTP') + + @patch("smtplib.SMTP") def test_send_password_reset_email(self, mock_smtp, email_service): """测试发送密码重置邮件""" mock_server = MagicMock() mock_smtp.return_value.__enter__.return_value = mock_server - + success, error = email_service.send_password_reset_email( to_email="user@example.com", username="TestUser", reset_url="https://example.com/reset?token=xyz789", ) - + assert success is True assert error is None - + # 验证发送了邮件 mock_server.sendmail.assert_called_once() call_args = mock_server.sendmail.call_args - + # 验证收件人 assert call_args[0][1] == ["user@example.com"] - - @patch('smtplib.SMTP') + + @patch("smtplib.SMTP") def test_send_workspace_invitation_email(self, mock_smtp, email_service): """测试发送工作空间邀请邮件""" mock_server = MagicMock() mock_smtp.return_value.__enter__.return_value = mock_server - + success, error = email_service.send_workspace_invitation_email( to_email="user@example.com", inviter_name="Alice", @@ -147,18 +150,18 @@ class TestEmailService: role="admin", invitation_url="https://example.com/invite?token=inv123", ) - + assert success is True assert error is None - + # 验证发送了邮件 mock_server.sendmail.assert_called_once() call_args = mock_server.sendmail.call_args - + # 验证收件人 assert call_args[0][1] == ["user@example.com"] - - @patch('smtplib.SMTP') + + @patch("smtplib.SMTP") def test_email_without_tls(self, mock_smtp): """测试不使用 TLS 发送邮件""" config = EmailConfig( @@ -170,27 +173,27 @@ class TestEmailService: use_tls=False, ) service = EmailService(config=config) - + mock_server = MagicMock() mock_smtp.return_value.__enter__.return_value = mock_server - + success, error = service.send_email( to_email="user@example.com", subject="Test", html_body="Test", ) - + assert success is True - + # 验证不调用 starttls mock_server.starttls.assert_not_called() # 验证不调用 login(没有用户名密码) mock_server.login.assert_not_called() - + def test_default_config(self): """测试默认配置""" service = EmailService() - + assert service.config.smtp_host == "smtp.gmail.com" assert service.config.smtp_port == 587 assert service.config.use_tls is True diff --git a/tests/unit/test_invite_member_use_case.py b/tests/unit/test_invite_member_use_case.py index 73ec5ee19..4dbc91539 100644 --- a/tests/unit/test_invite_member_use_case.py +++ b/tests/unit/test_invite_member_use_case.py @@ -1,52 +1,61 @@ """ 邀请成员 Use Case 测试 """ -import pytest -from unittest.mock import Mock + from datetime import datetime, timezone +from unittest.mock import Mock + +import pytest + from packages.application.workspace.invite_member_use_case import ( - InviteMemberUseCase, InviteMemberRequest, + InviteMemberUseCase, ) from packages.domain.entities import ( + User, Workspace, WorkspaceMember, WorkspaceMemberRole, - User, ) class TestInviteMemberUseCase: """邀请成员测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_workspace_and_user = Mock(return_value=None) return repo - + @pytest.fixture def mock_invitation_repo(self): repo = Mock() repo.find_pending_by_workspace_and_email = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def mock_user_repo(self): repo = Mock() repo.find_by_email = Mock(return_value=None) repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture - def use_case(self, mock_workspace_repo, mock_member_repo, mock_invitation_repo, mock_user_repo): + def use_case( + self, + mock_workspace_repo, + mock_member_repo, + mock_invitation_repo, + mock_user_repo, + ): email_service = Mock() email_service.send_workspace_invitation_email.return_value = (True, None) return InviteMemberUseCase( @@ -58,7 +67,7 @@ class TestInviteMemberUseCase: invitation_expire_days=7, email_service=email_service, ) - + @pytest.fixture def test_workspace(self): return Workspace( @@ -66,7 +75,7 @@ class TestInviteMemberUseCase: name="Test Workspace", owner_user_id="owner-id", ) - + @pytest.fixture def test_inviter(self): return User( @@ -75,7 +84,7 @@ class TestInviteMemberUseCase: username="inviter", display_name="Inviter User", ) - + @pytest.fixture def owner_member(self): return WorkspaceMember( @@ -84,7 +93,7 @@ class TestInviteMemberUseCase: user_id="inviter-id", role=WorkspaceMemberRole.OWNER, ) - + @pytest.fixture def admin_member(self): return WorkspaceMember( @@ -93,7 +102,7 @@ class TestInviteMemberUseCase: user_id="inviter-id", role=WorkspaceMemberRole.ADMIN, ) - + def test_invite_member_success_by_owner( self, use_case, @@ -109,31 +118,31 @@ class TestInviteMemberUseCase: mock_workspace_repo.find_by_id.return_value = test_workspace mock_member_repo.find_by_workspace_and_user.return_value = owner_member mock_user_repo.find_by_id.return_value = test_inviter - + request = InviteMemberRequest( workspace_id="workspace-123", inviter_user_id="inviter-id", invitee_email="newuser@test.com", role="member", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.invitee_email == "newuser@test.com" assert response.role == "member" assert response.expires_at is not None - + # 验证保存了邀请 mock_invitation_repo.save.assert_called_once() invitation = mock_invitation_repo.save.call_args[0][0] assert invitation.invitee_email == "newuser@test.com" assert invitation.status == "pending" - + # 验证发送了邮件 use_case.email_service.send_workspace_invitation_email.assert_called_once() - + def test_invite_member_success_by_admin( self, use_case, @@ -149,35 +158,35 @@ class TestInviteMemberUseCase: mock_workspace_repo.find_by_id.return_value = test_workspace mock_member_repo.find_by_workspace_and_user.return_value = admin_member mock_user_repo.find_by_id.return_value = test_inviter - + request = InviteMemberRequest( workspace_id="workspace-123", inviter_user_id="inviter-id", invitee_email="newuser@test.com", role="viewer", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None - + def test_invite_member_workspace_not_found(self, use_case, mock_workspace_repo): """测试 Workspace 不存在""" mock_workspace_repo.find_by_id.return_value = None - + request = InviteMemberRequest( workspace_id="nonexistent", inviter_user_id="inviter-id", invitee_email="newuser@test.com", role="member", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Workspace not found" - + def test_invite_member_inviter_not_member( self, use_case, @@ -188,19 +197,19 @@ class TestInviteMemberUseCase: """测试邀请人不是成员""" mock_workspace_repo.find_by_id.return_value = test_workspace mock_member_repo.find_by_workspace_and_user.return_value = None - + request = InviteMemberRequest( workspace_id="workspace-123", inviter_user_id="nonmember-id", invitee_email="newuser@test.com", role="member", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "You are not a member of this workspace" - + def test_invite_member_inviter_no_permission( self, use_case, @@ -210,7 +219,7 @@ class TestInviteMemberUseCase: ): """测试邀请人没有权限(只是普通成员)""" mock_workspace_repo.find_by_id.return_value = test_workspace - + regular_member = WorkspaceMember( id="member-3", workspace_id="workspace-123", @@ -218,19 +227,19 @@ class TestInviteMemberUseCase: role=WorkspaceMemberRole.MEMBER, ) mock_member_repo.find_by_workspace_and_user.return_value = regular_member - + request = InviteMemberRequest( workspace_id="workspace-123", inviter_user_id="inviter-id", invitee_email="newuser@test.com", role="member", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Only owners and admins can invite members" - + def test_invite_member_already_member( self, use_case, @@ -251,7 +260,7 @@ class TestInviteMemberUseCase: role=WorkspaceMemberRole.MEMBER, ), ] - + existing_user = User( id="existing-user-id", email="existing@test.com", @@ -259,19 +268,19 @@ class TestInviteMemberUseCase: display_name="Existing User", ) mock_user_repo.find_by_email.return_value = existing_user - + request = InviteMemberRequest( workspace_id="workspace-123", inviter_user_id="inviter-id", invitee_email="existing@test.com", role="member", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "User is already a member of this workspace" - + def test_invite_member_invalid_role(self, use_case): """测试无效角色(不能邀请为 owner)""" request = InviteMemberRequest( @@ -280,8 +289,8 @@ class TestInviteMemberUseCase: invitee_email="newuser@test.com", role="owner", ) - + response, error = use_case.execute(request) - + assert response is None assert "Invalid role" in error diff --git a/tests/unit/test_jwt_service.py b/tests/unit/test_jwt_service.py index 6a88fe263..7f9f1f7dd 100644 --- a/tests/unit/test_jwt_service.py +++ b/tests/unit/test_jwt_service.py @@ -1,135 +1,107 @@ """ JWT 工具类测试 """ -import pytest + from datetime import datetime, timedelta + +import pytest from jwt.exceptions import ExpiredSignatureError, InvalidTokenError -from packages.domain.auth.jwt_service import ( - JWTService, - JWTConfig, - TokenType, -) +from packages.domain.auth.jwt_service import JWTConfig, JWTService, TokenType class TestJWTService: """JWT 服务测试""" - + @pytest.fixture def jwt_service(self): """创建 JWT 服务实例""" config = JWTConfig() config.SECRET_KEY = "test-secret-key-for-testing" return JWTService(config) - + def test_create_access_token(self, jwt_service): """测试创建 access_token""" - token = jwt_service.create_access_token( - user_id="user-123", - workspace_id="workspace-456", - role="admin" - ) - + token = jwt_service.create_access_token(user_id="user-123", workspace_id="workspace-456", role="admin") + assert isinstance(token, str) assert len(token) > 0 - + # 验证 Token 内容 payload = jwt_service.verify_access_token(token) assert payload["sub"] == "user-123" assert payload["workspace_id"] == "workspace-456" assert payload["role"] == "admin" assert payload["type"] == TokenType.ACCESS - + def test_create_refresh_token(self, jwt_service): """测试创建 refresh_token""" - token = jwt_service.create_refresh_token( - user_id="user-123", - session_id="session-789" - ) - + token = jwt_service.create_refresh_token(user_id="user-123", session_id="session-789") + assert isinstance(token, str) assert len(token) > 0 - + # 验证 Token 内容 payload = jwt_service.verify_refresh_token(token) assert payload["sub"] == "user-123" assert payload["session_id"] == "session-789" assert payload["type"] == TokenType.REFRESH - + def test_verify_valid_access_token(self, jwt_service): """测试验证有效的 access_token""" - token = jwt_service.create_access_token( - user_id="user-123", - workspace_id="workspace-456", - role="member" - ) - + token = jwt_service.create_access_token(user_id="user-123", workspace_id="workspace-456", role="member") + payload = jwt_service.verify_access_token(token) assert payload["sub"] == "user-123" assert payload["workspace_id"] == "workspace-456" assert payload["role"] == "member" - + def test_verify_expired_token(self, jwt_service): """测试验证过期的 Token""" # 创建一个已过期的配置(使用相同的 SECRET_KEY) config = JWTConfig() config.SECRET_KEY = "test-secret-key-for-testing" # 与 fixture 相同 config.ACCESS_TOKEN_EXPIRE_MINUTES = -1 # 负数,立即过期 - + expired_service = JWTService(config) - token = expired_service.create_access_token( - user_id="user-123", - workspace_id="workspace-456", - role="admin" - ) - + token = expired_service.create_access_token(user_id="user-123", workspace_id="workspace-456", role="admin") + # 验证应该抛出过期异常 with pytest.raises(ExpiredSignatureError): jwt_service.verify_access_token(token) - + def test_verify_invalid_token(self, jwt_service): """测试验证无效的 Token""" invalid_token = "invalid.token.string" - + with pytest.raises(InvalidTokenError): jwt_service.verify_access_token(invalid_token) - + def test_verify_wrong_token_type(self, jwt_service): """测试验证错误类型的 Token""" # 创建 refresh_token - refresh_token = jwt_service.create_refresh_token( - user_id="user-123", - session_id="session-789" - ) - + refresh_token = jwt_service.create_refresh_token(user_id="user-123", session_id="session-789") + # 用 verify_access_token 验证应该失败 with pytest.raises(ValueError, match="Token type must be 'access'"): jwt_service.verify_access_token(refresh_token) - + # 反过来也一样 - access_token = jwt_service.create_access_token( - user_id="user-123", - workspace_id="workspace-456", - role="admin" - ) - + access_token = jwt_service.create_access_token(user_id="user-123", workspace_id="workspace-456", role="admin") + with pytest.raises(ValueError, match="Token type must be 'refresh'"): jwt_service.verify_refresh_token(access_token) - + def test_verify_tampered_token(self, jwt_service): """测试验证被篡改的 Token""" - token = jwt_service.create_access_token( - user_id="user-123", - workspace_id="workspace-456", - role="admin" - ) - + token = jwt_service.create_access_token(user_id="user-123", workspace_id="workspace-456", role="admin") + # 篡改 Token(修改最后几个字符) tampered_token = token[:-5] + "XXXXX" - + with pytest.raises(InvalidTokenError): jwt_service.verify_access_token(tampered_token) - + def test_additional_claims(self, jwt_service): """测试额外的声明""" token = jwt_service.create_access_token( @@ -138,27 +110,23 @@ class TestJWTService: role="admin", additional_claims={ "email": "user@example.com", - "display_name": "Test User" - } + "display_name": "Test User", + }, ) - + payload = jwt_service.verify_access_token(token) assert payload["email"] == "user@example.com" assert payload["display_name"] == "Test User" - + def test_decode_unsafe(self, jwt_service): """测试不安全解码(不验证签名)""" - token = jwt_service.create_access_token( - user_id="user-123", - workspace_id="workspace-456", - role="admin" - ) - + token = jwt_service.create_access_token(user_id="user-123", workspace_id="workspace-456", role="admin") + # 不验证签名地解码 payload = jwt_service.decode_token_unsafe(token) assert payload is not None assert payload["sub"] == "user-123" - + # 无效 Token 应该返回 None invalid_payload = jwt_service.decode_token_unsafe("invalid.token") assert invalid_payload is None diff --git a/tests/unit/test_list_members_use_case.py b/tests/unit/test_list_members_use_case.py index 7ffe4f66b..aeb33501f 100644 --- a/tests/unit/test_list_members_use_case.py +++ b/tests/unit/test_list_members_use_case.py @@ -1,43 +1,46 @@ """ 获取成员列表 Use Case 测试 """ -import pytest -from unittest.mock import Mock + from datetime import datetime, timezone +from unittest.mock import Mock + +import pytest + from packages.application.workspace.list_members_use_case import ( - ListMembersUseCase, ListMembersRequest, + ListMembersUseCase, ) from packages.domain.entities import ( + User, Workspace, WorkspaceMember, WorkspaceMemberRole, - User, ) class TestListMembersUseCase: """获取成员列表测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_workspace_and_user = Mock(return_value=None) repo.find_by_workspace = Mock(return_value=[]) return repo - + @pytest.fixture def mock_user_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture def use_case(self, mock_workspace_repo, mock_member_repo, mock_user_repo): return ListMembersUseCase( @@ -45,7 +48,7 @@ class TestListMembersUseCase: workspace_member_repository=mock_member_repo, user_repository=mock_user_repo, ) - + @pytest.fixture def test_workspace(self): return Workspace( @@ -53,7 +56,7 @@ class TestListMembersUseCase: name="Test Workspace", owner_user_id="owner-id", ) - + def test_list_members_success( self, use_case, @@ -64,7 +67,7 @@ class TestListMembersUseCase: ): """测试获取成员列表成功""" mock_workspace_repo.find_by_id.return_value = test_workspace - + # 请求者是 Admin requester_member = WorkspaceMember( id="member-1", @@ -73,7 +76,7 @@ class TestListMembersUseCase: role=WorkspaceMemberRole.ADMIN, ) mock_member_repo.find_by_workspace_and_user.return_value = requester_member - + # 3 个成员 member1 = WorkspaceMember( id="member-1", @@ -82,7 +85,7 @@ class TestListMembersUseCase: role=WorkspaceMemberRole.OWNER, invited_by=None, ) - + member2 = WorkspaceMember( id="member-2", workspace_id="workspace-123", @@ -90,7 +93,7 @@ class TestListMembersUseCase: role=WorkspaceMemberRole.ADMIN, invited_by="owner-id", ) - + member3 = WorkspaceMember( id="member-3", workspace_id="workspace-123", @@ -98,9 +101,9 @@ class TestListMembersUseCase: role=WorkspaceMemberRole.MEMBER, invited_by="admin-id", ) - + mock_member_repo.find_by_workspace.return_value = [member1, member2, member3] - + # 用户信息 user1 = User( id="owner-id", @@ -108,34 +111,34 @@ class TestListMembersUseCase: username="owner", display_name="Owner User", ) - + user2 = User( id="admin-id", email="admin@test.com", username="admin", display_name="Admin User", ) - + user3 = User( id="user-id", email="user@test.com", username="user", display_name="Regular User", ) - + mock_user_repo.find_by_id.side_effect = [user1, user2, user3] - + request = ListMembersRequest( workspace_id="workspace-123", requester_user_id="admin-id", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert len(response.members) == 3 - + # 验证第一个成员(Owner) m1 = response.members[0] assert m1.user_id == "owner-id" @@ -143,19 +146,19 @@ class TestListMembersUseCase: assert m1.email == "owner@test.com" assert m1.role == "owner" assert m1.invited_by is None - + # 验证第二个成员(Admin) m2 = response.members[1] assert m2.user_id == "admin-id" assert m2.role == "admin" assert m2.invited_by == "owner-id" - + # 验证第三个成员(Member) m3 = response.members[2] assert m3.user_id == "user-id" assert m3.role == "member" assert m3.invited_by == "admin-id" - + def test_list_members_workspace_not_found( self, use_case, @@ -163,17 +166,17 @@ class TestListMembersUseCase: ): """测试工作空间不存在""" mock_workspace_repo.find_by_id.return_value = None - + request = ListMembersRequest( workspace_id="nonexistent", requester_user_id="user-id", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Workspace not found" - + def test_list_members_requester_not_member( self, use_case, @@ -184,17 +187,17 @@ class TestListMembersUseCase: """测试请求者不是成员""" mock_workspace_repo.find_by_id.return_value = test_workspace mock_member_repo.find_by_workspace_and_user.return_value = None - + request = ListMembersRequest( workspace_id="workspace-123", requester_user_id="outsider-id", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "You are not a member of this workspace" - + def test_list_members_empty_workspace( self, use_case, @@ -204,7 +207,7 @@ class TestListMembersUseCase: ): """测试空工作空间(理论上不应该发生)""" mock_workspace_repo.find_by_id.return_value = test_workspace - + requester_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -213,38 +216,38 @@ class TestListMembersUseCase: ) mock_member_repo.find_by_workspace_and_user.return_value = requester_member mock_member_repo.find_by_workspace.return_value = [] - + request = ListMembersRequest( workspace_id="workspace-123", requester_user_id="user-id", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert len(response.members) == 0 - + def test_list_members_missing_workspace_id(self, use_case): """测试缺少工作空间 ID""" request = ListMembersRequest( workspace_id="", requester_user_id="user-id", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Workspace ID is required" - + def test_list_members_missing_requester_id(self, use_case): """测试缺少请求者 ID""" request = ListMembersRequest( workspace_id="workspace-123", requester_user_id="", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Requester user ID is required" diff --git a/tests/unit/test_list_workspaces_use_case.py b/tests/unit/test_list_workspaces_use_case.py index ccb2090e3..ca2334252 100644 --- a/tests/unit/test_list_workspaces_use_case.py +++ b/tests/unit/test_list_workspaces_use_case.py @@ -1,45 +1,44 @@ """ 获取工作空间列表和详情 Use Case 测试 """ -import pytest -from unittest.mock import Mock + from datetime import datetime, timezone +from unittest.mock import Mock + +import pytest + from packages.application.workspace.list_workspaces_use_case import ( - ListWorkspacesUseCase, - ListWorkspacesRequest, - GetWorkspaceDetailUseCase, GetWorkspaceDetailRequest, + GetWorkspaceDetailUseCase, + ListWorkspacesRequest, + ListWorkspacesUseCase, ) -from packages.domain.entities import ( - Workspace, - WorkspaceMember, - WorkspaceMemberRole, -) +from packages.domain.entities import Workspace, WorkspaceMember, WorkspaceMemberRole class TestListWorkspacesUseCase: """获取工作空间列表测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_user = Mock(return_value=[]) repo.count_by_workspace = Mock(return_value=0) return repo - + @pytest.fixture def use_case(self, mock_workspace_repo, mock_member_repo): return ListWorkspacesUseCase( workspace_repository=mock_workspace_repo, workspace_member_repository=mock_member_repo, ) - + def test_list_workspaces_success( self, use_case, @@ -54,40 +53,40 @@ class TestListWorkspacesUseCase: user_id="user-123", role=WorkspaceMemberRole.OWNER, ) - + membership2 = WorkspaceMember( id="member-2", workspace_id="workspace-2", user_id="user-123", role=WorkspaceMemberRole.MEMBER, ) - + mock_member_repo.find_by_user.return_value = [membership1, membership2] - + workspace1 = Workspace( id="workspace-1", name="My Workspace", owner_user_id="user-123", subscription_plan="free", ) - + workspace2 = Workspace( id="workspace-2", name="Team Workspace", owner_user_id="other-user", subscription_plan="pro", ) - + mock_workspace_repo.find_by_id.side_effect = [workspace1, workspace2] mock_member_repo.count_by_workspace.side_effect = [1, 5] - + request = ListWorkspacesRequest(user_id="user-123") response, error = use_case.execute(request) - + assert error is None assert response is not None assert len(response.workspaces) == 2 - + # 验证第一个工作空间 ws1 = response.workspaces[0] assert ws1.workspace_id == "workspace-1" @@ -95,7 +94,7 @@ class TestListWorkspacesUseCase: assert ws1.user_role == "owner" assert ws1.member_count == 1 assert ws1.subscription_plan == "free" - + # 验证第二个工作空间 ws2 = response.workspaces[1] assert ws2.workspace_id == "workspace-2" @@ -103,7 +102,7 @@ class TestListWorkspacesUseCase: assert ws2.user_role == "member" assert ws2.member_count == 5 assert ws2.subscription_plan == "pro" - + def test_list_workspaces_no_memberships( self, use_case, @@ -111,46 +110,46 @@ class TestListWorkspacesUseCase: ): """测试用户没有加入任何工作空间""" mock_member_repo.find_by_user.return_value = [] - + request = ListWorkspacesRequest(user_id="user-123") response, error = use_case.execute(request) - + assert error is None assert response is not None assert len(response.workspaces) == 0 - + def test_list_workspaces_missing_user_id(self, use_case): """测试缺少用户 ID""" request = ListWorkspacesRequest(user_id="") response, error = use_case.execute(request) - + assert response is None assert error == "User ID is required" class TestGetWorkspaceDetailUseCase: """获取工作空间详情测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_workspace_and_user = Mock(return_value=None) repo.count_by_workspace = Mock(return_value=0) return repo - + @pytest.fixture def use_case(self, mock_workspace_repo, mock_member_repo): return GetWorkspaceDetailUseCase( workspace_repository=mock_workspace_repo, workspace_member_repository=mock_member_repo, ) - + @pytest.fixture def test_workspace(self): return Workspace( @@ -163,7 +162,7 @@ class TestGetWorkspaceDetailUseCase: max_storage_gb=100, used_storage_gb=25.5, ) - + def test_get_workspace_detail_success( self, use_case, @@ -173,7 +172,7 @@ class TestGetWorkspaceDetailUseCase: ): """测试获取工作空间详情成功""" mock_workspace_repo.find_by_id.return_value = test_workspace - + membership = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -182,14 +181,14 @@ class TestGetWorkspaceDetailUseCase: ) mock_member_repo.find_by_workspace_and_user.return_value = membership mock_member_repo.count_by_workspace.return_value = 8 - + request = GetWorkspaceDetailRequest( workspace_id="workspace-123", user_id="user-123", ) - + detail, error = use_case.execute(request) - + assert error is None assert detail is not None assert detail.workspace_id == "workspace-123" @@ -201,7 +200,7 @@ class TestGetWorkspaceDetailUseCase: assert detail.used_storage_gb == 25.5 assert detail.member_count == 8 assert detail.user_role == "admin" - + def test_get_workspace_detail_not_found( self, use_case, @@ -209,17 +208,17 @@ class TestGetWorkspaceDetailUseCase: ): """测试工作空间不存在""" mock_workspace_repo.find_by_id.return_value = None - + request = GetWorkspaceDetailRequest( workspace_id="nonexistent", user_id="user-123", ) - + detail, error = use_case.execute(request) - + assert detail is None assert error == "Workspace not found" - + def test_get_workspace_detail_not_member( self, use_case, @@ -230,37 +229,37 @@ class TestGetWorkspaceDetailUseCase: """测试用户不是成员""" mock_workspace_repo.find_by_id.return_value = test_workspace mock_member_repo.find_by_workspace_and_user.return_value = None - + request = GetWorkspaceDetailRequest( workspace_id="workspace-123", user_id="user-123", ) - + detail, error = use_case.execute(request) - + assert detail is None assert error == "You are not a member of this workspace" - + def test_get_workspace_detail_missing_workspace_id(self, use_case): """测试缺少工作空间 ID""" request = GetWorkspaceDetailRequest( workspace_id="", user_id="user-123", ) - + detail, error = use_case.execute(request) - + assert detail is None assert error == "Workspace ID is required" - + def test_get_workspace_detail_missing_user_id(self, use_case): """测试缺少用户 ID""" request = GetWorkspaceDetailRequest( workspace_id="workspace-123", user_id="", ) - + detail, error = use_case.execute(request) - + assert detail is None assert error == "User ID is required" diff --git a/tests/unit/test_login_use_case.py b/tests/unit/test_login_use_case.py index 642149ee8..6cd0cc923 100644 --- a/tests/unit/test_login_use_case.py +++ b/tests/unit/test_login_use_case.py @@ -1,35 +1,38 @@ """ 用户登录 Use Case 测试 """ -import pytest -from unittest.mock import Mock + from datetime import datetime, timezone +from unittest.mock import Mock + +import pytest + from packages.application.auth import ( - LoginUseCase, LoginRequest, - LogoutUseCase, + LoginUseCase, LogoutRequest, + LogoutUseCase, ) -from packages.domain.entities import User from packages.domain.auth import password_hasher +from packages.domain.entities import User class TestLoginUseCase: """登录用例测试""" - + @pytest.fixture def mock_user_repo(self): repo = Mock() repo.find_by_email = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def use_case(self, mock_user_repo): session_store = Mock() session_store.save_session.return_value = True return LoginUseCase(user_repository=mock_user_repo, session_store=session_store) - + @pytest.fixture def test_user(self): """创建测试用户""" @@ -42,20 +45,20 @@ class TestLoginUseCase: password_hash=password_hash, email_verified=True, ) - + def test_login_success(self, use_case, mock_user_repo, test_user): """测试登录成功""" mock_user_repo.find_by_email.return_value = test_user - + request = LoginRequest( email="test@example.com", password="SecurePass123", device_info="Chrome/Windows", ip_address="192.168.1.1", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.user_id == "user-123" @@ -64,124 +67,124 @@ class TestLoginUseCase: assert response.access_token != "" assert response.refresh_token != "" assert response.expires_in > 0 - + # 验证保存了 session use_case.session_store.save_session.assert_called_once() - + # 验证更新了最后登录信息 mock_user_repo.save.assert_called_once() - + def test_login_invalid_email(self, use_case, mock_user_repo): """测试邮箱不存在""" mock_user_repo.find_by_email.return_value = None - + request = LoginRequest( email="nonexistent@example.com", password="SecurePass123", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Invalid email or password" - + def test_login_wrong_password(self, use_case, mock_user_repo, test_user): """测试密码错误""" mock_user_repo.find_by_email.return_value = test_user - + request = LoginRequest( email="test@example.com", password="WrongPassword123", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Invalid email or password" - + def test_login_missing_email(self, use_case): """测试缺少邮箱""" request = LoginRequest( email="", password="SecurePass123", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Email is required" - + def test_login_missing_password(self, use_case, mock_user_repo, test_user): """测试缺少密码""" mock_user_repo.find_by_email.return_value = test_user - + request = LoginRequest( email="test@example.com", password="", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Password is required" class TestLogoutUseCase: """登出用例测试""" - + @pytest.fixture def use_case(self): session_store = Mock() return LogoutUseCase(session_store=session_store) - + def test_logout_current_device(self, use_case): """测试登出当前设备""" use_case.session_store.delete_session.return_value = True - + request = LogoutRequest( user_id="user-123", session_id="session-abc", logout_all_devices=False, ) - + success, error = use_case.execute(request) - + assert success is True assert error is None - + use_case.session_store.delete_session.assert_called_once_with("session-abc") - + def test_logout_all_devices(self, use_case): """测试登出所有设备""" use_case.session_store.delete_all_user_sessions.return_value = 3 - + request = LogoutRequest( user_id="user-123", logout_all_devices=True, ) - + success, error = use_case.execute(request) - + assert success is True assert error is None - + use_case.session_store.delete_all_user_sessions.assert_called_once_with("user-123") - + def test_logout_session_not_found(self, use_case): """测试 session 不存在""" use_case.session_store.delete_session.return_value = False - + request = LogoutRequest( user_id="user-123", session_id="nonexistent", logout_all_devices=False, ) - + success, error = use_case.execute(request) - + assert success is False assert error == "Session not found" - + def test_logout_missing_session_id(self, use_case): """测试缺少 session_id""" request = LogoutRequest( @@ -189,8 +192,8 @@ class TestLogoutUseCase: session_id=None, logout_all_devices=False, ) - + success, error = use_case.execute(request) - + assert success is False assert error == "Session ID is required" diff --git a/tests/unit/test_password_hasher.py b/tests/unit/test_password_hasher.py index eee3509ec..93921d85f 100644 --- a/tests/unit/test_password_hasher.py +++ b/tests/unit/test_password_hasher.py @@ -1,92 +1,91 @@ """ 密码哈希工具测试 """ + import pytest -from packages.domain.auth.password_hasher import ( - PasswordHasher, - PasswordValidator, -) + +from packages.domain.auth.password_hasher import PasswordHasher, PasswordValidator class TestPasswordHasher: """密码哈希测试""" - + @pytest.fixture def hasher(self): """创建密码哈希器""" return PasswordHasher(rounds=4) # 测试用低 cost,加快速度 - + def test_hash_password(self, hasher): """测试密码哈希""" password = "MySecurePassword123" hashed = hasher.hash_password(password) - + assert isinstance(hashed, str) assert len(hashed) > 0 assert hashed != password # 哈希后不等于原文 assert hashed.startswith("$2b$") # bcrypt 格式 - + def test_hash_same_password_different_result(self, hasher): """测试相同密码每次哈希结果不同(因为 salt 不同)""" password = "MySecurePassword123" hash1 = hasher.hash_password(password) hash2 = hasher.hash_password(password) - + assert hash1 != hash2 # salt 不同,哈希不同 - + def test_verify_correct_password(self, hasher): """测试验证正确的密码""" password = "MySecurePassword123" hashed = hasher.hash_password(password) - + assert hasher.verify_password(password, hashed) is True - + def test_verify_incorrect_password(self, hasher): """测试验证错误的密码""" password = "MySecurePassword123" hashed = hasher.hash_password(password) - + assert hasher.verify_password("WrongPassword", hashed) is False - + def test_verify_empty_password(self, hasher): """测试空密码验证""" hashed = hasher.hash_password("test") - + assert hasher.verify_password("", hashed) is False - + def test_verify_empty_hash(self, hasher): """测试空哈希验证""" assert hasher.verify_password("test", "") is False - + def test_verify_invalid_hash(self, hasher): """测试无效的哈希""" assert hasher.verify_password("test", "invalid-hash") is False - + def test_hash_empty_password(self, hasher): """测试哈希空密码应该失败""" with pytest.raises(ValueError, match="Password cannot be empty"): hasher.hash_password("") - + def test_invalid_rounds(self): """测试无效的 rounds 参数""" with pytest.raises(ValueError, match="rounds must be between 4 and 31"): PasswordHasher(rounds=2) - + with pytest.raises(ValueError, match="rounds must be between 4 and 31"): PasswordHasher(rounds=50) - + def test_unicode_password(self, hasher): """测试 Unicode 密码""" password = "密码123!@#" hashed = hasher.hash_password(password) - + assert hasher.verify_password(password, hashed) is True assert hasher.verify_password("错误密码", hashed) is False class TestPasswordValidator: """密码验证器测试""" - + @pytest.fixture def validator(self): """创建密码验证器""" @@ -97,37 +96,37 @@ class TestPasswordValidator: require_digit=True, require_special=False, ) - + def test_valid_password(self, validator): """测试有效密码""" valid, error = validator.validate("MyPassword123") assert valid is True assert error is None - + def test_password_too_short(self, validator): """测试密码太短""" valid, error = validator.validate("Pass1") assert valid is False assert "at least 8 characters" in error - + def test_password_no_uppercase(self, validator): """测试没有大写字母""" valid, error = validator.validate("mypassword123") assert valid is False assert "uppercase letter" in error - + def test_password_no_lowercase(self, validator): """测试没有小写字母""" valid, error = validator.validate("MYPASSWORD123") assert valid is False assert "lowercase letter" in error - + def test_password_no_digit(self, validator): """测试没有数字""" valid, error = validator.validate("MyPassword") assert valid is False assert "digit" in error - + def test_password_with_special_chars(self): """测试要求特殊字符""" validator = PasswordValidator( @@ -137,23 +136,23 @@ class TestPasswordValidator: require_digit=True, require_special=True, ) - + # 没有特殊字符 valid, error = validator.validate("MyPassword123") assert valid is False assert "special character" in error - + # 有特殊字符 valid, error = validator.validate("MyPassword123!") assert valid is True assert error is None - + def test_empty_password(self, validator): """测试空密码""" valid, error = validator.validate("") assert valid is False assert "cannot be empty" in error - + def test_custom_min_length(self): """测试自定义最小长度""" validator = PasswordValidator( @@ -163,11 +162,11 @@ class TestPasswordValidator: require_digit=False, require_special=False, ) - + valid, error = validator.validate("short") assert valid is False assert "at least 12 characters" in error - + valid, error = validator.validate("longenoughpassword") assert valid is True assert error is None diff --git a/tests/unit/test_password_reset_use_case.py b/tests/unit/test_password_reset_use_case.py index 762ed9413..ced1f05b0 100644 --- a/tests/unit/test_password_reset_use_case.py +++ b/tests/unit/test_password_reset_use_case.py @@ -1,28 +1,31 @@ """ 密码重置 Use Case 测试 """ -import pytest -from unittest.mock import Mock + from datetime import datetime, timedelta, timezone +from unittest.mock import Mock + +import pytest + from packages.application.auth.password_reset_use_case import ( - RequestPasswordResetUseCase, RequestPasswordResetRequest, - ResetPasswordUseCase, + RequestPasswordResetUseCase, ResetPasswordRequest, + ResetPasswordUseCase, ) from packages.domain.entities import User class TestRequestPasswordResetUseCase: """请求密码重置测试""" - + @pytest.fixture def mock_user_repo(self): repo = Mock() repo.find_by_email = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def use_case(self, mock_user_repo): email_service = Mock() @@ -33,7 +36,7 @@ class TestRequestPasswordResetUseCase: token_expire_hours=1, email_service=email_service, ) - + @pytest.fixture def test_user(self): return User( @@ -43,62 +46,62 @@ class TestRequestPasswordResetUseCase: display_name="Test User", password_hash="hash", ) - + def test_request_reset_success(self, use_case, mock_user_repo, test_user): """测试请求重置成功""" mock_user_repo.find_by_email.return_value = test_user - + request = RequestPasswordResetRequest(email="test@example.com") success, error = use_case.execute(request) - + assert success is True assert error is None - + # 验证保存了用户 mock_user_repo.save.assert_called_once() saved_user = mock_user_repo.save.call_args[0][0] assert saved_user.password_reset_token is not None assert saved_user.password_reset_expires_at is not None - + # 验证发送了邮件 use_case.email_service.send_password_reset_email.assert_called_once() - + def test_request_reset_user_not_exists(self, use_case, mock_user_repo): """测试用户不存在(仍返回成功,避免暴露)""" mock_user_repo.find_by_email.return_value = None - + request = RequestPasswordResetRequest(email="nonexistent@example.com") success, error = use_case.execute(request) - + assert success is True # 安全考虑,仍返回成功 assert error is None - + # 不发送邮件 use_case.email_service.send_password_reset_email.assert_not_called() - + def test_request_reset_missing_email(self, use_case): """测试缺少邮箱""" request = RequestPasswordResetRequest(email="") success, error = use_case.execute(request) - + assert success is False assert error == "Email is required" class TestResetPasswordUseCase: """重置密码测试""" - + @pytest.fixture def mock_user_repo(self): repo = Mock() repo.find_by_password_reset_token = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def use_case(self, mock_user_repo): return ResetPasswordUseCase(user_repository=mock_user_repo) - + @pytest.fixture def test_user(self): return User( @@ -110,68 +113,68 @@ class TestResetPasswordUseCase: password_reset_token="valid-token", password_reset_expires_at=datetime.now(timezone.utc) + timedelta(hours=1), ) - + def test_reset_password_success(self, use_case, mock_user_repo, test_user): """测试重置密码成功""" mock_user_repo.find_by_password_reset_token.return_value = test_user - + request = ResetPasswordRequest( token="valid-token", new_password="NewSecurePass123", ) success, error = use_case.execute(request) - + assert success is True assert error is None - + # 验证密码已更新 assert test_user.password_hash != "old-hash" assert test_user.password_reset_token is None assert test_user.password_reset_expires_at is None - + # 验证保存了用户 mock_user_repo.save.assert_called_once() - + def test_reset_password_weak_password(self, use_case, mock_user_repo, test_user): """测试弱密码""" mock_user_repo.find_by_password_reset_token.return_value = test_user - + request = ResetPasswordRequest( token="valid-token", new_password="weak", ) success, error = use_case.execute(request) - + assert success is False assert "at least 8 characters" in error - + def test_reset_password_invalid_token(self, use_case, mock_user_repo): """测试无效令牌""" mock_user_repo.find_by_password_reset_token.return_value = None - + request = ResetPasswordRequest( token="invalid-token", new_password="NewSecurePass123", ) success, error = use_case.execute(request) - + assert success is False assert error == "Invalid or expired reset token" - + def test_reset_password_expired_token(self, use_case, mock_user_repo, test_user): """测试过期令牌""" test_user.password_reset_expires_at = datetime.now(timezone.utc) - timedelta(hours=1) mock_user_repo.find_by_password_reset_token.return_value = test_user - + request = ResetPasswordRequest( token="valid-token", new_password="NewSecurePass123", ) success, error = use_case.execute(request) - + assert success is False assert error == "Reset token has expired" - + def test_reset_password_missing_token(self, use_case): """测试缺少令牌""" request = ResetPasswordRequest( @@ -179,19 +182,19 @@ class TestResetPasswordUseCase: new_password="NewSecurePass123", ) success, error = use_case.execute(request) - + assert success is False assert error == "Reset token is required" - + def test_reset_password_missing_password(self, use_case, mock_user_repo, test_user): """测试缺少新密码""" mock_user_repo.find_by_password_reset_token.return_value = test_user - + request = ResetPasswordRequest( token="valid-token", new_password="", ) success, error = use_case.execute(request) - + assert success is False assert error == "New password is required" diff --git a/tests/unit/test_permissions.py b/tests/unit/test_permissions.py index a841fbb3c..df86fdaff 100644 --- a/tests/unit/test_permissions.py +++ b/tests/unit/test_permissions.py @@ -1,32 +1,28 @@ """ 权限验证辅助函数测试 """ -import pytest + from unittest.mock import Mock -from packages.domain.permissions import ( - PermissionChecker, - Permission, - has_permission, -) -from packages.domain.entities import ( - WorkspaceMember, - WorkspaceMemberRole, -) + +import pytest + +from packages.domain.entities import WorkspaceMember, WorkspaceMemberRole +from packages.domain.permissions import Permission, PermissionChecker, has_permission class TestPermissionChecker: """权限检查器测试""" - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_workspace_and_user = Mock(return_value=None) return repo - + @pytest.fixture def checker(self, mock_member_repo): return PermissionChecker(workspace_member_repository=mock_member_repo) - + def test_check_workspace_access_has_access(self, checker, mock_member_repo): """测试有访问权限""" member = WorkspaceMember( @@ -36,21 +32,21 @@ class TestPermissionChecker: role=WorkspaceMemberRole.MEMBER, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + has_access, role = checker.check_workspace_access("workspace-123", "user-123") - + assert has_access is True assert role == "member" - + def test_check_workspace_access_no_access(self, checker, mock_member_repo): """测试无访问权限""" mock_member_repo.find_by_workspace_and_user.return_value = None - + has_access, role = checker.check_workspace_access("workspace-123", "user-123") - + assert has_access is False assert role is None - + def test_check_is_owner_true(self, checker, mock_member_repo): """测试是 Owner""" member = WorkspaceMember( @@ -60,11 +56,11 @@ class TestPermissionChecker: role=WorkspaceMemberRole.OWNER, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + is_owner = checker.check_is_owner("workspace-123", "user-123") - + assert is_owner is True - + def test_check_is_owner_false(self, checker, mock_member_repo): """测试不是 Owner""" member = WorkspaceMember( @@ -74,11 +70,11 @@ class TestPermissionChecker: role=WorkspaceMemberRole.ADMIN, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + is_owner = checker.check_is_owner("workspace-123", "user-123") - + assert is_owner is False - + def test_check_is_admin_or_owner_admin(self, checker, mock_member_repo): """测试是 Admin""" member = WorkspaceMember( @@ -88,11 +84,11 @@ class TestPermissionChecker: role=WorkspaceMemberRole.ADMIN, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + is_admin_or_owner = checker.check_is_admin_or_owner("workspace-123", "user-123") - + assert is_admin_or_owner is True - + def test_check_is_admin_or_owner_owner(self, checker, mock_member_repo): """测试是 Owner""" member = WorkspaceMember( @@ -102,11 +98,11 @@ class TestPermissionChecker: role=WorkspaceMemberRole.OWNER, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + is_admin_or_owner = checker.check_is_admin_or_owner("workspace-123", "user-123") - + assert is_admin_or_owner is True - + def test_check_is_admin_or_owner_member(self, checker, mock_member_repo): """测试是普通成员""" member = WorkspaceMember( @@ -116,11 +112,11 @@ class TestPermissionChecker: role=WorkspaceMemberRole.MEMBER, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + is_admin_or_owner = checker.check_is_admin_or_owner("workspace-123", "user-123") - + assert is_admin_or_owner is False - + def test_check_can_create_project_member(self, checker, mock_member_repo): """测试 Member 可以创建项目""" member = WorkspaceMember( @@ -130,11 +126,11 @@ class TestPermissionChecker: role=WorkspaceMemberRole.MEMBER, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + can_create = checker.check_can_create_project("workspace-123", "user-123") - + assert can_create is True - + def test_check_can_create_project_viewer(self, checker, mock_member_repo): """测试 Viewer 不能创建项目""" member = WorkspaceMember( @@ -144,11 +140,11 @@ class TestPermissionChecker: role=WorkspaceMemberRole.VIEWER, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + can_create = checker.check_can_create_project("workspace-123", "user-123") - + assert can_create is False - + def test_check_can_delete_project_member(self, checker, mock_member_repo): """测试 Member 不能删除项目""" member = WorkspaceMember( @@ -158,11 +154,11 @@ class TestPermissionChecker: role=WorkspaceMemberRole.MEMBER, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + can_delete = checker.check_can_delete_project("workspace-123", "user-123") - + assert can_delete is False - + def test_check_can_delete_project_admin(self, checker, mock_member_repo): """测试 Admin 可以删除项目""" member = WorkspaceMember( @@ -172,33 +168,33 @@ class TestPermissionChecker: role=WorkspaceMemberRole.ADMIN, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + can_delete = checker.check_can_delete_project("workspace-123", "user-123") - + assert can_delete is True class TestPermissionFunctions: """权限函数测试""" - + def test_has_permission_owner(self): """测试 Owner 权限""" assert has_permission(WorkspaceMemberRole.OWNER, Permission.WORKSPACE_DELETE) is True assert has_permission(WorkspaceMemberRole.OWNER, Permission.MEMBER_REMOVE) is True assert has_permission(WorkspaceMemberRole.OWNER, Permission.PROJECT_CREATE) is True - + def test_has_permission_admin(self): """测试 Admin 权限""" assert has_permission(WorkspaceMemberRole.ADMIN, Permission.WORKSPACE_EDIT) is True assert has_permission(WorkspaceMemberRole.ADMIN, Permission.MEMBER_REMOVE) is True assert has_permission(WorkspaceMemberRole.ADMIN, Permission.WORKSPACE_DELETE) is False - + def test_has_permission_member(self): """测试 Member 权限""" assert has_permission(WorkspaceMemberRole.MEMBER, Permission.PROJECT_CREATE) is True assert has_permission(WorkspaceMemberRole.MEMBER, Permission.PROJECT_DELETE) is False assert has_permission(WorkspaceMemberRole.MEMBER, Permission.MEMBER_INVITE) is False - + def test_has_permission_viewer(self): """测试 Viewer 权限""" assert has_permission(WorkspaceMemberRole.VIEWER, Permission.WORKSPACE_VIEW) is True diff --git a/tests/unit/test_quota.py b/tests/unit/test_quota.py index 41650eaac..f5f7eceb5 100644 --- a/tests/unit/test_quota.py +++ b/tests/unit/test_quota.py @@ -1,39 +1,38 @@ """ 配额检查服务测试 """ -import pytest + from unittest.mock import Mock -from packages.domain.quota import ( - QuotaChecker, - QuotaWarningLevel, - get_warning_level, -) + +import pytest + from packages.domain.entities import Workspace +from packages.domain.quota import QuotaChecker, QuotaWarningLevel, get_warning_level class TestQuotaChecker: """配额检查器测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def mock_project_repo(self): repo = Mock() repo.count_by_workspace = Mock(return_value=0) return repo - + @pytest.fixture def checker(self, mock_workspace_repo, mock_project_repo): return QuotaChecker( workspace_repository=mock_workspace_repo, project_repository=mock_project_repo, ) - + def test_check_can_create_project_within_limit( self, checker, @@ -50,12 +49,12 @@ class TestQuotaChecker: ) mock_workspace_repo.find_by_id.return_value = workspace mock_project_repo.count_by_workspace.return_value = 2 - + can_create, error = checker.check_can_create_project("workspace-123") - + assert can_create is True assert error is None - + def test_check_can_create_project_at_limit( self, checker, @@ -72,12 +71,12 @@ class TestQuotaChecker: ) mock_workspace_repo.find_by_id.return_value = workspace mock_project_repo.count_by_workspace.return_value = 3 - + can_create, error = checker.check_can_create_project("workspace-123") - + assert can_create is False assert "Project limit reached" in error - + def test_check_can_create_project_unlimited( self, checker, @@ -93,13 +92,13 @@ class TestQuotaChecker: max_projects=999999, ) mock_workspace_repo.find_by_id.return_value = workspace - mock_project_repo.count_by_workspace.return_value=1000 - + mock_project_repo.count_by_workspace.return_value = 1000 + can_create, error = checker.check_can_create_project("workspace-123") - + assert can_create is True assert error is None - + def test_check_storage_available_within_limit( self, checker, @@ -115,12 +114,12 @@ class TestQuotaChecker: used_storage_gb=5.0, ) mock_workspace_repo.find_by_id.return_value = workspace - + can_store, error = checker.check_storage_available("workspace-123", 3.0) - + assert can_store is True assert error is None - + def test_check_storage_available_exceeded( self, checker, @@ -136,12 +135,12 @@ class TestQuotaChecker: used_storage_gb=8.0, ) mock_workspace_repo.find_by_id.return_value = workspace - + can_store, error = checker.check_storage_available("workspace-123", 3.0) - + assert can_store is False assert "Storage limit exceeded" in error - + def test_get_quota_status( self, checker, @@ -160,9 +159,9 @@ class TestQuotaChecker: ) mock_workspace_repo.find_by_id.return_value = workspace mock_project_repo.count_by_workspace.return_value = 2 - + status = checker.get_quota_status("workspace-123") - + assert status is not None assert status["workspace_id"] == "workspace-123" assert status["subscription_plan"] == "free" @@ -173,7 +172,7 @@ class TestQuotaChecker: assert status["storage"]["limit_gb"] == 10 assert status["storage"]["remaining_gb"] == 2.5 assert status["storage"]["usage_percent"] == 75.0 - + def test_get_quota_status_unlimited( self, checker, @@ -192,12 +191,12 @@ class TestQuotaChecker: ) mock_workspace_repo.find_by_id.return_value = workspace mock_project_repo.count_by_workspace.return_value = 1000 - + status = checker.get_quota_status("workspace-123") - + assert status["projects"]["unlimited"] is True assert status["projects"]["usage_percent"] == 0 - + def test_update_storage_usage_increase( self, checker, @@ -212,14 +211,14 @@ class TestQuotaChecker: used_storage_gb=5.0, ) mock_workspace_repo.find_by_id.return_value = workspace - + success, error = checker.update_storage_usage("workspace-123", 2.5) - + assert success is True assert error is None assert workspace.used_storage_gb == 7.5 mock_workspace_repo.save.assert_called_once() - + def test_update_storage_usage_decrease( self, checker, @@ -234,12 +233,12 @@ class TestQuotaChecker: used_storage_gb=5.0, ) mock_workspace_repo.find_by_id.return_value = workspace - + success, error = checker.update_storage_usage("workspace-123", -2.0) - + assert success is True assert workspace.used_storage_gb == 3.0 - + def test_update_storage_usage_prevent_negative( self, checker, @@ -254,33 +253,33 @@ class TestQuotaChecker: used_storage_gb=2.0, ) mock_workspace_repo.find_by_id.return_value = workspace - + success, error = checker.update_storage_usage("workspace-123", -5.0) - + assert success is True assert workspace.used_storage_gb == 0.0 class TestWarningLevel: """警告级别测试""" - + def test_get_warning_level_normal(self): """测试正常级别""" assert get_warning_level(50.0) == QuotaWarningLevel.NORMAL assert get_warning_level(79.9) == QuotaWarningLevel.NORMAL - + def test_get_warning_level_warning(self): """测试警告级别""" assert get_warning_level(80.0) == QuotaWarningLevel.WARNING assert get_warning_level(85.0) == QuotaWarningLevel.WARNING assert get_warning_level(89.9) == QuotaWarningLevel.WARNING - + def test_get_warning_level_critical(self): """测试严重级别""" assert get_warning_level(90.0) == QuotaWarningLevel.CRITICAL assert get_warning_level(95.0) == QuotaWarningLevel.CRITICAL assert get_warning_level(99.9) == QuotaWarningLevel.CRITICAL - + def test_get_warning_level_exceeded(self): """测试超出级别""" assert get_warning_level(100.0) == QuotaWarningLevel.EXCEEDED diff --git a/tests/unit/test_register_user_use_case.py b/tests/unit/test_register_user_use_case.py index 4cfa8405b..d82dda303 100644 --- a/tests/unit/test_register_user_use_case.py +++ b/tests/unit/test_register_user_use_case.py @@ -1,20 +1,23 @@ """ 用户注册 Use Case 测试 """ -import pytest + from unittest.mock import Mock + +import pytest + from packages.application.auth import ( - RegisterUserUseCase, RegisterUserRequest, - VerifyEmailUseCase, + RegisterUserUseCase, VerifyEmailRequest, + VerifyEmailUseCase, ) from packages.domain.entities import User class TestRegisterUserUseCase: """注册用例测试""" - + @pytest.fixture def mock_user_repo(self): """Mock 用户仓储""" @@ -24,7 +27,7 @@ class TestRegisterUserUseCase: repo.find_by_verification_token = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def use_case(self, mock_user_repo): """创建注册用例""" @@ -35,26 +38,26 @@ class TestRegisterUserUseCase: base_url="https://test.com", email_service=email_service, ) - + def test_register_user_success(self, use_case, mock_user_repo): """测试注册成功""" - + request = RegisterUserRequest( email="test@example.com", password="SecurePass123", username="testuser", display_name="Test User", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.email == "test@example.com" assert response.username == "testuser" assert response.display_name == "Test User" assert response.email_verification_sent is True - + # 验证保存了用户 mock_user_repo.save.assert_called_once() saved_user = mock_user_repo.save.call_args[0][0] @@ -62,7 +65,7 @@ class TestRegisterUserUseCase: assert saved_user.password_hash != "" assert saved_user.email_verified is False assert saved_user.email_verification_token is not None - + def test_register_user_weak_password(self, use_case): """测试弱密码""" request = RegisterUserRequest( @@ -71,13 +74,13 @@ class TestRegisterUserUseCase: username="testuser", display_name="Test User", ) - + response, error = use_case.execute(request) - + assert response is None assert error is not None assert "at least 8 characters" in error - + def test_register_user_email_exists(self, use_case, mock_user_repo): """测试邮箱已存在""" # Mock 返回已存在的用户 @@ -88,19 +91,19 @@ class TestRegisterUserUseCase: display_name="Existing", ) mock_user_repo.find_by_email.return_value = existing_user - + request = RegisterUserRequest( email="test@example.com", password="SecurePass123", username="testuser", display_name="Test User", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Email already registered" - + def test_register_user_username_taken(self, use_case, mock_user_repo): """测试用户名已被占用""" existing_user = User( @@ -110,19 +113,19 @@ class TestRegisterUserUseCase: display_name="Other", ) mock_user_repo.find_by_username.return_value = existing_user - + request = RegisterUserRequest( email="test@example.com", password="SecurePass123", username="testuser", display_name="Test User", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Username already taken" - + def test_register_user_missing_email(self, use_case): """测试缺少邮箱""" request = RegisterUserRequest( @@ -131,25 +134,28 @@ class TestRegisterUserUseCase: username="testuser", display_name="Test User", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Email is required" - + def test_register_user_email_send_failure(self, use_case, mock_user_repo): """测试邮件发送失败(用户仍然创建)""" - use_case.email_service.send_verification_email.return_value = (False, "SMTP error") - + use_case.email_service.send_verification_email.return_value = ( + False, + "SMTP error", + ) + request = RegisterUserRequest( email="test@example.com", password="SecurePass123", username="testuser", display_name="Test User", ) - + response, error = use_case.execute(request) - + assert error is None # 用户创建成功 assert response is not None assert response.email_verification_sent is False # 但邮件发送失败 @@ -157,18 +163,18 @@ class TestRegisterUserUseCase: class TestVerifyEmailUseCase: """邮箱验证用例测试""" - + @pytest.fixture def mock_user_repo(self): repo = Mock() repo.find_by_verification_token = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def use_case(self, mock_user_repo): return VerifyEmailUseCase(user_repository=mock_user_repo) - + def test_verify_email_success(self, use_case, mock_user_repo): """测试验证成功""" user = User( @@ -180,28 +186,28 @@ class TestVerifyEmailUseCase: email_verification_token="valid-token", ) mock_user_repo.find_by_verification_token.return_value = user - + request = VerifyEmailRequest(token="valid-token") success, error = use_case.execute(request) - + assert success is True assert error is None - + # 验证用户状态已更新 assert user.email_verified is True assert user.email_verification_token is None mock_user_repo.save.assert_called_once() - + def test_verify_email_invalid_token(self, use_case, mock_user_repo): """测试无效令牌""" mock_user_repo.find_by_verification_token.return_value = None - + request = VerifyEmailRequest(token="invalid-token") success, error = use_case.execute(request) - + assert success is False assert error == "Invalid or expired verification token" - + def test_verify_email_already_verified(self, use_case, mock_user_repo): """测试已验证的邮箱""" user = User( @@ -213,9 +219,9 @@ class TestVerifyEmailUseCase: email_verification_token="old-token", ) mock_user_repo.find_by_verification_token.return_value = user - + request = VerifyEmailRequest(token="old-token") success, error = use_case.execute(request) - + assert success is True # 已验证也返回成功 assert error is None diff --git a/tests/unit/test_remove_member_use_case.py b/tests/unit/test_remove_member_use_case.py index fed2cb084..555b4fdc6 100644 --- a/tests/unit/test_remove_member_use_case.py +++ b/tests/unit/test_remove_member_use_case.py @@ -1,44 +1,43 @@ """ 移除成员 Use Case 测试 """ -import pytest + from unittest.mock import Mock + +import pytest + from packages.application.workspace.remove_member_use_case import ( - RemoveMemberUseCase, - RemoveMemberRequest, - LeaveWorkspaceUseCase, LeaveWorkspaceRequest, + LeaveWorkspaceUseCase, + RemoveMemberRequest, + RemoveMemberUseCase, ) -from packages.domain.entities import ( - Workspace, - WorkspaceMember, - WorkspaceMemberRole, -) +from packages.domain.entities import Workspace, WorkspaceMember, WorkspaceMemberRole class TestRemoveMemberUseCase: """移除成员测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_workspace_and_user = Mock(return_value=None) repo.delete = Mock(return_value=True) return repo - + @pytest.fixture def use_case(self, mock_workspace_repo, mock_member_repo): return RemoveMemberUseCase( workspace_repository=mock_workspace_repo, workspace_member_repository=mock_member_repo, ) - + @pytest.fixture def test_workspace(self): return Workspace( @@ -46,7 +45,7 @@ class TestRemoveMemberUseCase: name="Test Workspace", owner_user_id="owner-id", ) - + def test_remove_member_success_by_owner( self, use_case, @@ -56,35 +55,38 @@ class TestRemoveMemberUseCase: ): """测试 Owner 移除成员成功""" mock_workspace_repo.find_by_id.return_value = test_workspace - + owner_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="owner-id", role=WorkspaceMemberRole.OWNER, ) - + target_member = WorkspaceMember( id="member-2", workspace_id="workspace-123", user_id="target-id", role=WorkspaceMemberRole.MEMBER, ) - - mock_member_repo.find_by_workspace_and_user.side_effect = [owner_member, target_member] - + + mock_member_repo.find_by_workspace_and_user.side_effect = [ + owner_member, + target_member, + ] + request = RemoveMemberRequest( workspace_id="workspace-123", requester_user_id="owner-id", target_user_id="target-id", ) - + success, error = use_case.execute(request) - + assert success is True assert error is None mock_member_repo.delete.assert_called_once_with("member-2") - + def test_remove_member_success_by_admin( self, use_case, @@ -94,34 +96,37 @@ class TestRemoveMemberUseCase: ): """测试 Admin 移除普通成员成功""" mock_workspace_repo.find_by_id.return_value = test_workspace - + admin_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="admin-id", role=WorkspaceMemberRole.ADMIN, ) - + target_member = WorkspaceMember( id="member-2", workspace_id="workspace-123", user_id="target-id", role=WorkspaceMemberRole.VIEWER, ) - - mock_member_repo.find_by_workspace_and_user.side_effect = [admin_member, target_member] - + + mock_member_repo.find_by_workspace_and_user.side_effect = [ + admin_member, + target_member, + ] + request = RemoveMemberRequest( workspace_id="workspace-123", requester_user_id="admin-id", target_user_id="target-id", ) - + success, error = use_case.execute(request) - + assert success is True assert error is None - + def test_remove_member_cannot_remove_owner( self, use_case, @@ -131,34 +136,37 @@ class TestRemoveMemberUseCase: ): """测试不能移除 Owner""" mock_workspace_repo.find_by_id.return_value = test_workspace - + admin_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="admin-id", role=WorkspaceMemberRole.ADMIN, ) - + owner_member = WorkspaceMember( id="member-2", workspace_id="workspace-123", user_id="owner-id", role=WorkspaceMemberRole.OWNER, ) - - mock_member_repo.find_by_workspace_and_user.side_effect = [admin_member, owner_member] - + + mock_member_repo.find_by_workspace_and_user.side_effect = [ + admin_member, + owner_member, + ] + request = RemoveMemberRequest( workspace_id="workspace-123", requester_user_id="admin-id", target_user_id="owner-id", ) - + success, error = use_case.execute(request) - + assert success is False assert error == "Cannot remove the workspace owner" - + def test_remove_member_admin_cannot_remove_admin( self, use_case, @@ -168,34 +176,37 @@ class TestRemoveMemberUseCase: ): """测试 Admin 不能移除另一个 Admin""" mock_workspace_repo.find_by_id.return_value = test_workspace - + admin_member1 = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="admin-id-1", role=WorkspaceMemberRole.ADMIN, ) - + admin_member2 = WorkspaceMember( id="member-2", workspace_id="workspace-123", user_id="admin-id-2", role=WorkspaceMemberRole.ADMIN, ) - - mock_member_repo.find_by_workspace_and_user.side_effect = [admin_member1, admin_member2] - + + mock_member_repo.find_by_workspace_and_user.side_effect = [ + admin_member1, + admin_member2, + ] + request = RemoveMemberRequest( workspace_id="workspace-123", requester_user_id="admin-id-1", target_user_id="admin-id-2", ) - + success, error = use_case.execute(request) - + assert success is False assert error == "Admins cannot remove other admins" - + def test_remove_member_cannot_remove_self( self, use_case, @@ -205,27 +216,27 @@ class TestRemoveMemberUseCase: ): """测试不能移除自己""" mock_workspace_repo.find_by_id.return_value = test_workspace - + admin_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="admin-id", role=WorkspaceMemberRole.ADMIN, ) - + mock_member_repo.find_by_workspace_and_user.return_value = admin_member - + request = RemoveMemberRequest( workspace_id="workspace-123", requester_user_id="admin-id", target_user_id="admin-id", ) - + success, error = use_case.execute(request) - + assert success is False assert "Cannot remove yourself" in error - + def test_remove_member_no_permission( self, use_case, @@ -235,51 +246,51 @@ class TestRemoveMemberUseCase: ): """测试普通成员没有权限移除""" mock_workspace_repo.find_by_id.return_value = test_workspace - + regular_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="user-id", role=WorkspaceMemberRole.MEMBER, ) - + mock_member_repo.find_by_workspace_and_user.return_value = regular_member - + request = RemoveMemberRequest( workspace_id="workspace-123", requester_user_id="user-id", target_user_id="target-id", ) - + success, error = use_case.execute(request) - + assert success is False assert error == "Only owners and admins can remove members" class TestLeaveWorkspaceUseCase: """离开 Workspace 测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_workspace_and_user = Mock(return_value=None) repo.delete = Mock(return_value=True) return repo - + @pytest.fixture def use_case(self, mock_workspace_repo, mock_member_repo): return LeaveWorkspaceUseCase( workspace_repository=mock_workspace_repo, workspace_member_repository=mock_member_repo, ) - + @pytest.fixture def test_workspace(self): return Workspace( @@ -287,7 +298,7 @@ class TestLeaveWorkspaceUseCase: name="Test Workspace", owner_user_id="owner-id", ) - + def test_leave_workspace_success( self, use_case, @@ -297,7 +308,7 @@ class TestLeaveWorkspaceUseCase: ): """测试离开工作空间成功""" mock_workspace_repo.find_by_id.return_value = test_workspace - + member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -305,18 +316,18 @@ class TestLeaveWorkspaceUseCase: role=WorkspaceMemberRole.MEMBER, ) mock_member_repo.find_by_workspace_and_user.return_value = member - + request = LeaveWorkspaceRequest( workspace_id="workspace-123", user_id="user-id", ) - + success, error = use_case.execute(request) - + assert success is True assert error is None mock_member_repo.delete.assert_called_once_with("member-1") - + def test_leave_workspace_owner_cannot_leave( self, use_case, @@ -326,7 +337,7 @@ class TestLeaveWorkspaceUseCase: ): """测试 Owner 不能离开""" mock_workspace_repo.find_by_id.return_value = test_workspace - + owner_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -334,17 +345,17 @@ class TestLeaveWorkspaceUseCase: role=WorkspaceMemberRole.OWNER, ) mock_member_repo.find_by_workspace_and_user.return_value = owner_member - + request = LeaveWorkspaceRequest( workspace_id="workspace-123", user_id="owner-id", ) - + success, error = use_case.execute(request) - + assert success is False assert "Owner cannot leave workspace" in error - + def test_leave_workspace_not_member( self, use_case, @@ -355,13 +366,13 @@ class TestLeaveWorkspaceUseCase: """测试不是成员""" mock_workspace_repo.find_by_id.return_value = test_workspace mock_member_repo.find_by_workspace_and_user.return_value = None - + request = LeaveWorkspaceRequest( workspace_id="workspace-123", user_id="user-id", ) - + success, error = use_case.execute(request) - + assert success is False assert error == "You are not a member of this workspace" diff --git a/tests/unit/test_session_store.py b/tests/unit/test_session_store.py index 05f971dd3..44c8392a2 100644 --- a/tests/unit/test_session_store.py +++ b/tests/unit/test_session_store.py @@ -1,17 +1,19 @@ """ Redis Session 存储测试 """ -import pytest + import json from datetime import datetime -from unittest.mock import Mock, MagicMock +from unittest.mock import MagicMock, Mock + +import pytest from packages.domain.auth.session_store import SessionStore class TestSessionStore: """Session 存储测试""" - + @pytest.fixture def mock_redis(self): """创建 Mock Redis 客户端""" @@ -19,46 +21,46 @@ class TestSessionStore: redis_mock.data = {} # 模拟内存存储 redis_mock.expires = {} # 模拟过期时间 redis_mock.sets = {} # 模拟集合 - + def setex(key, seconds, value): redis_mock.data[key] = value redis_mock.expires[key] = seconds return True - + def get(key): return redis_mock.data.get(key) - + def delete(key): if key in redis_mock.data: del redis_mock.data[key] return 1 return 0 - + def exists(key): return 1 if key in redis_mock.data else 0 - + def ttl(key): return redis_mock.expires.get(key, -1) - + def sadd(key, *values): if key not in redis_mock.sets: redis_mock.sets[key] = set() redis_mock.sets[key].update(values) return len(values) - + def smembers(key): return redis_mock.sets.get(key, set()) - + def srem(key, *values): if key in redis_mock.sets: redis_mock.sets[key].discard(*values) return len(values) return 0 - + def expire(key, seconds): redis_mock.expires[key] = seconds return True - + redis_mock.setex = setex redis_mock.get = get redis_mock.delete = delete @@ -68,14 +70,14 @@ class TestSessionStore: redis_mock.smembers = smembers redis_mock.srem = srem redis_mock.expire = expire - + return redis_mock - + @pytest.fixture def session_store(self, mock_redis): """创建 Session 存储实例""" return SessionStore(redis_client=mock_redis) - + def test_save_session(self, session_store, mock_redis): """测试保存 Session""" result = session_store.save_session( @@ -86,27 +88,27 @@ class TestSessionStore: ip_address="192.168.1.1", expires_in_seconds=3600, ) - + assert result is True - + # 验证数据已保存 session_key = "session:session-123" assert session_key in mock_redis.data - + session_data = json.loads(mock_redis.data[session_key]) assert session_data["session_id"] == "session-123" assert session_data["user_id"] == "user-456" assert session_data["device_info"] == "Chrome/Windows" assert session_data["ip_address"] == "192.168.1.1" - + # 验证 refresh_token 已保存 refresh_token_key = "refresh_token:session-123" assert mock_redis.data[refresh_token_key] == "refresh-token-abc" - + # 验证用户 Session 集合已更新 user_sessions_key = "user_sessions:user-456" assert "session-123" in mock_redis.sets[user_sessions_key] - + def test_get_session(self, session_store, mock_redis): """测试获取 Session""" # 先保存 @@ -117,20 +119,20 @@ class TestSessionStore: device_info="Chrome", ip_address="127.0.0.1", ) - + # 获取 session = session_store.get_session("session-123") - + assert session is not None assert session["session_id"] == "session-123" assert session["user_id"] == "user-456" assert session["device_info"] == "Chrome" - + def test_get_nonexistent_session(self, session_store): """测试获取不存在的 Session""" session = session_store.get_session("nonexistent") assert session is None - + def test_get_refresh_token(self, session_store): """测试获取 refresh_token""" session_store.save_session( @@ -140,10 +142,10 @@ class TestSessionStore: device_info="Chrome", ip_address="127.0.0.1", ) - + token = session_store.get_refresh_token("session-123") assert token == "my-refresh-token" - + def test_update_last_active(self, session_store, mock_redis): """测试更新最后活跃时间""" session_store.save_session( @@ -153,19 +155,19 @@ class TestSessionStore: device_info="Chrome", ip_address="127.0.0.1", ) - + # 获取原始时间 session1 = session_store.get_session("session-123") original_time = session1["last_active_at"] - + # 更新 result = session_store.update_last_active("session-123") assert result is True - + # 验证时间已更新 session2 = session_store.get_session("session-123") assert session2["last_active_at"] >= original_time - + def test_delete_session(self, session_store, mock_redis): """测试删除 Session""" session_store.save_session( @@ -175,22 +177,22 @@ class TestSessionStore: device_info="Chrome", ip_address="127.0.0.1", ) - + # 删除 result = session_store.delete_session("session-123") assert result is True - + # 验证已删除 session = session_store.get_session("session-123") assert session is None - + token = session_store.get_refresh_token("session-123") assert token is None - + # 验证从用户集合中移除 user_sessions_key = "user_sessions:user-456" assert "session-123" not in mock_redis.sets.get(user_sessions_key, set()) - + def test_get_user_sessions(self, session_store): """测试获取用户的所有 Session""" # 创建多个 Session @@ -201,7 +203,7 @@ class TestSessionStore: device_info="Chrome", ip_address="192.168.1.1", ) - + session_store.save_session( session_id="session-2", user_id="user-456", @@ -209,15 +211,15 @@ class TestSessionStore: device_info="Firefox", ip_address="192.168.1.2", ) - + # 获取 sessions = session_store.get_user_sessions("user-456") - + assert len(sessions) == 2 session_ids = [s["session_id"] for s in sessions] assert "session-1" in session_ids assert "session-2" in session_ids - + def test_delete_all_user_sessions(self, session_store): """测试删除用户的所有 Session""" # 创建多个 Session @@ -228,7 +230,7 @@ class TestSessionStore: device_info="Chrome", ip_address="127.0.0.1", ) - + session_store.save_session( session_id="session-2", user_id="user-456", @@ -236,19 +238,19 @@ class TestSessionStore: device_info="Firefox", ip_address="127.0.0.1", ) - + # 删除所有 count = session_store.delete_all_user_sessions("user-456") assert count == 2 - + # 验证已删除 sessions = session_store.get_user_sessions("user-456") assert len(sessions) == 0 - + def test_session_exists(self, session_store): """测试检查 Session 是否存在""" assert session_store.session_exists("nonexistent") is False - + session_store.save_session( session_id="session-123", user_id="user-456", @@ -256,9 +258,9 @@ class TestSessionStore: device_info="Chrome", ip_address="127.0.0.1", ) - + assert session_store.session_exists("session-123") is True - + def test_multiple_users(self, session_store): """测试多用户隔离""" # 用户 1 的 Session @@ -269,7 +271,7 @@ class TestSessionStore: device_info="Chrome", ip_address="127.0.0.1", ) - + # 用户 2 的 Session session_store.save_session( session_id="session-user2", @@ -278,12 +280,12 @@ class TestSessionStore: device_info="Firefox", ip_address="127.0.0.1", ) - + # 验证隔离 user1_sessions = session_store.get_user_sessions("user-1") assert len(user1_sessions) == 1 assert user1_sessions[0]["session_id"] == "session-user1" - + user2_sessions = session_store.get_user_sessions("user-2") assert len(user2_sessions) == 1 assert user2_sessions[0]["session_id"] == "session-user2" diff --git a/tests/unit/test_subscription_use_case.py b/tests/unit/test_subscription_use_case.py index 532dc5d14..77fae08e4 100644 --- a/tests/unit/test_subscription_use_case.py +++ b/tests/unit/test_subscription_use_case.py @@ -1,44 +1,43 @@ """ Subscription 管理 Use Case 测试 """ -import pytest + from unittest.mock import Mock + +import pytest + from packages.application.workspace.subscription_use_case import ( - UpgradeSubscriptionUseCase, - UpgradeSubscriptionRequest, - CancelSubscriptionUseCase, CancelSubscriptionRequest, + CancelSubscriptionUseCase, + UpgradeSubscriptionRequest, + UpgradeSubscriptionUseCase, ) -from packages.domain.entities import ( - Workspace, - WorkspaceMember, - WorkspaceMemberRole, -) +from packages.domain.entities import Workspace, WorkspaceMember, WorkspaceMemberRole class TestUpgradeSubscriptionUseCase: """升级订阅测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_workspace_and_user = Mock(return_value=None) return repo - + @pytest.fixture def use_case(self, mock_workspace_repo, mock_member_repo): return UpgradeSubscriptionUseCase( workspace_repository=mock_workspace_repo, workspace_member_repository=mock_member_repo, ) - + def test_upgrade_from_free_to_pro( self, use_case, @@ -53,7 +52,7 @@ class TestUpgradeSubscriptionUseCase: subscription_plan="free", ) mock_workspace_repo.find_by_id.return_value = workspace - + owner_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -61,27 +60,27 @@ class TestUpgradeSubscriptionUseCase: role=WorkspaceMemberRole.OWNER, ) mock_member_repo.find_by_workspace_and_user.return_value = owner_member - + request = UpgradeSubscriptionRequest( workspace_id="workspace-123", requester_user_id="owner-id", new_plan="pro", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.old_plan == "free" assert response.new_plan == "pro" assert response.max_projects == 999999 assert response.max_storage_gb == 100 - + # 验证更新了 workspace assert workspace.subscription_plan == "pro" assert workspace.max_projects == 999999 assert workspace.subscription_expires_at is not None - + def test_upgrade_from_pro_to_enterprise( self, use_case, @@ -98,7 +97,7 @@ class TestUpgradeSubscriptionUseCase: max_storage_gb=100, ) mock_workspace_repo.find_by_id.return_value = workspace - + owner_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -106,21 +105,21 @@ class TestUpgradeSubscriptionUseCase: role=WorkspaceMemberRole.OWNER, ) mock_member_repo.find_by_workspace_and_user.return_value = owner_member - + request = UpgradeSubscriptionRequest( workspace_id="workspace-123", requester_user_id="owner-id", new_plan="enterprise", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.old_plan == "pro" assert response.new_plan == "enterprise" assert response.max_storage_gb == 1000 - + def test_upgrade_cannot_downgrade( self, use_case, @@ -135,7 +134,7 @@ class TestUpgradeSubscriptionUseCase: subscription_plan="pro", ) mock_workspace_repo.find_by_id.return_value = workspace - + owner_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -143,18 +142,18 @@ class TestUpgradeSubscriptionUseCase: role=WorkspaceMemberRole.OWNER, ) mock_member_repo.find_by_workspace_and_user.return_value = owner_member - + request = UpgradeSubscriptionRequest( workspace_id="workspace-123", requester_user_id="owner-id", new_plan="free", ) - + response, error = use_case.execute(request) - + assert response is None assert "Cannot downgrade" in error - + def test_upgrade_already_on_plan( self, use_case, @@ -169,7 +168,7 @@ class TestUpgradeSubscriptionUseCase: subscription_plan="pro", ) mock_workspace_repo.find_by_id.return_value = workspace - + owner_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -177,18 +176,18 @@ class TestUpgradeSubscriptionUseCase: role=WorkspaceMemberRole.OWNER, ) mock_member_repo.find_by_workspace_and_user.return_value = owner_member - + request = UpgradeSubscriptionRequest( workspace_id="workspace-123", requester_user_id="owner-id", new_plan="pro", ) - + response, error = use_case.execute(request) - + assert response is None assert "already on pro plan" in error - + def test_upgrade_only_owner_can_upgrade( self, use_case, @@ -203,7 +202,7 @@ class TestUpgradeSubscriptionUseCase: subscription_plan="free", ) mock_workspace_repo.find_by_id.return_value = workspace - + admin_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -211,42 +210,42 @@ class TestUpgradeSubscriptionUseCase: role=WorkspaceMemberRole.ADMIN, ) mock_member_repo.find_by_workspace_and_user.return_value = admin_member - + request = UpgradeSubscriptionRequest( workspace_id="workspace-123", requester_user_id="admin-id", new_plan="pro", ) - + response, error = use_case.execute(request) - + assert response is None assert "Only workspace owner" in error class TestCancelSubscriptionUseCase: """取消订阅测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_workspace_and_user = Mock(return_value=None) return repo - + @pytest.fixture def use_case(self, mock_workspace_repo, mock_member_repo): return CancelSubscriptionUseCase( workspace_repository=mock_workspace_repo, workspace_member_repository=mock_member_repo, ) - + def test_cancel_subscription_success( self, use_case, @@ -263,7 +262,7 @@ class TestCancelSubscriptionUseCase: max_storage_gb=100, ) mock_workspace_repo.find_by_id.return_value = workspace - + owner_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -271,23 +270,23 @@ class TestCancelSubscriptionUseCase: role=WorkspaceMemberRole.OWNER, ) mock_member_repo.find_by_workspace_and_user.return_value = owner_member - + request = CancelSubscriptionRequest( workspace_id="workspace-123", requester_user_id="owner-id", ) - + success, error = use_case.execute(request) - + assert success is True assert error is None - + # 验证降级到 free assert workspace.subscription_plan == "free" assert workspace.max_projects == 3 assert workspace.max_storage_gb == 10 assert workspace.subscription_expires_at is None - + def test_cancel_already_free( self, use_case, @@ -302,7 +301,7 @@ class TestCancelSubscriptionUseCase: subscription_plan="free", ) mock_workspace_repo.find_by_id.return_value = workspace - + owner_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -310,17 +309,17 @@ class TestCancelSubscriptionUseCase: role=WorkspaceMemberRole.OWNER, ) mock_member_repo.find_by_workspace_and_user.return_value = owner_member - + request = CancelSubscriptionRequest( workspace_id="workspace-123", requester_user_id="owner-id", ) - + success, error = use_case.execute(request) - + assert success is False assert "already on free plan" in error - + def test_cancel_only_owner_can_cancel( self, use_case, @@ -335,7 +334,7 @@ class TestCancelSubscriptionUseCase: subscription_plan="pro", ) mock_workspace_repo.find_by_id.return_value = workspace - + admin_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", @@ -343,13 +342,13 @@ class TestCancelSubscriptionUseCase: role=WorkspaceMemberRole.ADMIN, ) mock_member_repo.find_by_workspace_and_user.return_value = admin_member - + request = CancelSubscriptionRequest( workspace_id="workspace-123", requester_user_id="admin-id", ) - + success, error = use_case.execute(request) - + assert success is False assert "Only workspace owner" in error diff --git a/tests/unit/test_update_member_role_use_case.py b/tests/unit/test_update_member_role_use_case.py index 68673427d..134bae1e5 100644 --- a/tests/unit/test_update_member_role_use_case.py +++ b/tests/unit/test_update_member_role_use_case.py @@ -1,42 +1,41 @@ """ 修改成员角色 Use Case 测试 """ -import pytest + from unittest.mock import Mock + +import pytest + from packages.application.workspace.update_member_role_use_case import ( - UpdateMemberRoleUseCase, UpdateMemberRoleRequest, + UpdateMemberRoleUseCase, ) -from packages.domain.entities import ( - Workspace, - WorkspaceMember, - WorkspaceMemberRole, -) +from packages.domain.entities import Workspace, WorkspaceMember, WorkspaceMemberRole class TestUpdateMemberRoleUseCase: """修改成员角色测试""" - + @pytest.fixture def mock_workspace_repo(self): repo = Mock() repo.find_by_id = Mock(return_value=None) return repo - + @pytest.fixture def mock_member_repo(self): repo = Mock() repo.find_by_workspace_and_user = Mock(return_value=None) repo.save = Mock() return repo - + @pytest.fixture def use_case(self, mock_workspace_repo, mock_member_repo): return UpdateMemberRoleUseCase( workspace_repository=mock_workspace_repo, workspace_member_repository=mock_member_repo, ) - + @pytest.fixture def test_workspace(self): return Workspace( @@ -44,7 +43,7 @@ class TestUpdateMemberRoleUseCase: name="Test Workspace", owner_user_id="owner-id", ) - + def test_update_role_success_by_owner( self, use_case, @@ -54,42 +53,45 @@ class TestUpdateMemberRoleUseCase: ): """测试 Owner 修改成员角色成功""" mock_workspace_repo.find_by_id.return_value = test_workspace - + owner_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="owner-id", role=WorkspaceMemberRole.OWNER, ) - + target_member = WorkspaceMember( id="member-2", workspace_id="workspace-123", user_id="target-id", role=WorkspaceMemberRole.MEMBER, ) - - mock_member_repo.find_by_workspace_and_user.side_effect = [owner_member, target_member] - + + mock_member_repo.find_by_workspace_and_user.side_effect = [ + owner_member, + target_member, + ] + request = UpdateMemberRoleRequest( workspace_id="workspace-123", requester_user_id="owner-id", target_user_id="target-id", new_role="admin", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None assert response.user_id == "target-id" assert response.old_role == "member" assert response.new_role == "admin" - + # 验证更新了角色 assert target_member.role == "admin" mock_member_repo.save.assert_called_once() - + def test_update_role_success_by_admin( self, use_case, @@ -99,35 +101,38 @@ class TestUpdateMemberRoleUseCase: ): """测试 Admin 修改普通成员角色成功""" mock_workspace_repo.find_by_id.return_value = test_workspace - + admin_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="admin-id", role=WorkspaceMemberRole.ADMIN, ) - + target_member = WorkspaceMember( id="member-2", workspace_id="workspace-123", user_id="target-id", role=WorkspaceMemberRole.VIEWER, ) - - mock_member_repo.find_by_workspace_and_user.side_effect = [admin_member, target_member] - + + mock_member_repo.find_by_workspace_and_user.side_effect = [ + admin_member, + target_member, + ] + request = UpdateMemberRoleRequest( workspace_id="workspace-123", requester_user_id="admin-id", target_user_id="target-id", new_role="member", ) - + response, error = use_case.execute(request) - + assert error is None assert response is not None - + def test_update_role_cannot_change_owner( self, use_case, @@ -137,35 +142,38 @@ class TestUpdateMemberRoleUseCase: ): """测试不能修改 Owner 角色""" mock_workspace_repo.find_by_id.return_value = test_workspace - + admin_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="admin-id", role=WorkspaceMemberRole.ADMIN, ) - + owner_member = WorkspaceMember( id="member-2", workspace_id="workspace-123", user_id="owner-id", role=WorkspaceMemberRole.OWNER, ) - - mock_member_repo.find_by_workspace_and_user.side_effect = [admin_member, owner_member] - + + mock_member_repo.find_by_workspace_and_user.side_effect = [ + admin_member, + owner_member, + ] + request = UpdateMemberRoleRequest( workspace_id="workspace-123", requester_user_id="admin-id", target_user_id="owner-id", new_role="member", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Cannot change the owner's role" - + def test_update_role_admin_cannot_change_admin( self, use_case, @@ -175,35 +183,38 @@ class TestUpdateMemberRoleUseCase: ): """测试 Admin 不能修改另一个 Admin 角色""" mock_workspace_repo.find_by_id.return_value = test_workspace - + admin_member1 = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="admin-id-1", role=WorkspaceMemberRole.ADMIN, ) - + admin_member2 = WorkspaceMember( id="member-2", workspace_id="workspace-123", user_id="admin-id-2", role=WorkspaceMemberRole.ADMIN, ) - - mock_member_repo.find_by_workspace_and_user.side_effect = [admin_member1, admin_member2] - + + mock_member_repo.find_by_workspace_and_user.side_effect = [ + admin_member1, + admin_member2, + ] + request = UpdateMemberRoleRequest( workspace_id="workspace-123", requester_user_id="admin-id-1", target_user_id="admin-id-2", new_role="member", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Admins cannot change other admins' roles" - + def test_update_role_cannot_change_self( self, use_case, @@ -213,28 +224,28 @@ class TestUpdateMemberRoleUseCase: ): """测试不能修改自己的角色""" mock_workspace_repo.find_by_id.return_value = test_workspace - + admin_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="admin-id", role=WorkspaceMemberRole.ADMIN, ) - + mock_member_repo.find_by_workspace_and_user.return_value = admin_member - + request = UpdateMemberRoleRequest( workspace_id="workspace-123", requester_user_id="admin-id", target_user_id="admin-id", new_role="member", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Cannot change your own role" - + def test_update_role_invalid_role(self, use_case): """测试无效角色(不能改为 owner)""" request = UpdateMemberRoleRequest( @@ -243,12 +254,12 @@ class TestUpdateMemberRoleUseCase: target_user_id="target-id", new_role="owner", ) - + response, error = use_case.execute(request) - + assert response is None assert "Invalid role" in error - + def test_update_role_already_has_role( self, use_case, @@ -258,35 +269,38 @@ class TestUpdateMemberRoleUseCase: ): """测试角色相同""" mock_workspace_repo.find_by_id.return_value = test_workspace - + owner_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="owner-id", role=WorkspaceMemberRole.OWNER, ) - + target_member = WorkspaceMember( id="member-2", workspace_id="workspace-123", user_id="target-id", role=WorkspaceMemberRole.ADMIN, ) - - mock_member_repo.find_by_workspace_and_user.side_effect = [owner_member, target_member] - + + mock_member_repo.find_by_workspace_and_user.side_effect = [ + owner_member, + target_member, + ] + request = UpdateMemberRoleRequest( workspace_id="workspace-123", requester_user_id="owner-id", target_user_id="target-id", new_role="admin", ) - + response, error = use_case.execute(request) - + assert response is None assert "already has the admin role" in error - + def test_update_role_no_permission( self, use_case, @@ -296,24 +310,24 @@ class TestUpdateMemberRoleUseCase: ): """测试普通成员没有权限""" mock_workspace_repo.find_by_id.return_value = test_workspace - + regular_member = WorkspaceMember( id="member-1", workspace_id="workspace-123", user_id="user-id", role=WorkspaceMemberRole.MEMBER, ) - + mock_member_repo.find_by_workspace_and_user.return_value = regular_member - + request = UpdateMemberRoleRequest( workspace_id="workspace-123", requester_user_id="user-id", target_user_id="target-id", new_role="admin", ) - + response, error = use_case.execute(request) - + assert response is None assert error == "Only owners and admins can change member roles"