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