style: normalize python formatting gates

This commit is contained in:
Xiaoxia AI
2026-06-21 06:52:19 +08:00
parent 0809a079c5
commit bfbaddbd9a
129 changed files with 3024 additions and 2485 deletions
+1 -2
View File
@@ -1,8 +1,7 @@
import os
from logging.config import fileConfig
from sqlalchemy import engine_from_config
from sqlalchemy import pool
from sqlalchemy import engine_from_config, pool
from alembic import context
+85 -15
View File
@@ -12,9 +12,9 @@ run this migration normally.
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
from alembic import op
revision: str = "001"
down_revision: Union[str, None] = None
@@ -67,8 +67,18 @@ def upgrade() -> None:
sa.PrimaryKeyConstraint("id"),
)
op.create_index(op.f("ix_asset_libraries_kind"), "asset_libraries", ["kind"], unique=False)
op.create_index(op.f("ix_asset_libraries_project_id"), "asset_libraries", ["project_id"], unique=False)
op.create_index(op.f("ix_asset_libraries_workspace_id"), "asset_libraries", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_asset_libraries_project_id"),
"asset_libraries",
["project_id"],
unique=False,
)
op.create_index(
op.f("ix_asset_libraries_workspace_id"),
"asset_libraries",
["workspace_id"],
unique=False,
)
op.create_table(
"assets",
@@ -96,7 +106,12 @@ def upgrade() -> None:
sa.PrimaryKeyConstraint("id"),
)
op.create_index(op.f("ix_assets_asset_library_id"), "assets", ["asset_library_id"], unique=False)
op.create_index(op.f("ix_assets_classification_status"), "assets", ["classification_status"], unique=False)
op.create_index(
op.f("ix_assets_classification_status"),
"assets",
["classification_status"],
unique=False,
)
op.create_index(op.f("ix_assets_created_at"), "assets", ["created_at"], unique=False)
op.create_index(op.f("ix_assets_file_type"), "assets", ["file_type"], unique=False)
op.create_index(op.f("ix_assets_project_id"), "assets", ["project_id"], unique=False)
@@ -119,7 +134,12 @@ def upgrade() -> None:
)
op.create_index(op.f("ix_ingest_jobs_library_id"), "ingest_jobs", ["library_id"], unique=False)
op.create_index(op.f("ix_ingest_jobs_project_id"), "ingest_jobs", ["project_id"], unique=False)
op.create_index(op.f("ix_ingest_jobs_workspace_id"), "ingest_jobs", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_ingest_jobs_workspace_id"),
"ingest_jobs",
["workspace_id"],
unique=False,
)
op.create_table(
"classification_jobs",
@@ -135,9 +155,24 @@ def upgrade() -> None:
sa.Column("updated_at", sa.DateTime(), nullable=False),
sa.PrimaryKeyConstraint("id"),
)
op.create_index(op.f("ix_classification_jobs_asset_id"), "classification_jobs", ["asset_id"], unique=False)
op.create_index(op.f("ix_classification_jobs_project_id"), "classification_jobs", ["project_id"], unique=False)
op.create_index(op.f("ix_classification_jobs_workspace_id"), "classification_jobs", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_classification_jobs_asset_id"),
"classification_jobs",
["asset_id"],
unique=False,
)
op.create_index(
op.f("ix_classification_jobs_project_id"),
"classification_jobs",
["project_id"],
unique=False,
)
op.create_index(
op.f("ix_classification_jobs_workspace_id"),
"classification_jobs",
["workspace_id"],
unique=False,
)
op.create_table(
"generation_tasks",
@@ -157,10 +192,25 @@ def upgrade() -> None:
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.PrimaryKeyConstraint("id"),
)
op.create_index(op.f("ix_generation_tasks_asset_library_id"), "generation_tasks", ["asset_library_id"], unique=False)
op.create_index(op.f("ix_generation_tasks_project_id"), "generation_tasks", ["project_id"], unique=False)
op.create_index(
op.f("ix_generation_tasks_asset_library_id"),
"generation_tasks",
["asset_library_id"],
unique=False,
)
op.create_index(
op.f("ix_generation_tasks_project_id"),
"generation_tasks",
["project_id"],
unique=False,
)
op.create_index(op.f("ix_generation_tasks_status"), "generation_tasks", ["status"], unique=False)
op.create_index(op.f("ix_generation_tasks_workspace_id"), "generation_tasks", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_generation_tasks_workspace_id"),
"generation_tasks",
["workspace_id"],
unique=False,
)
op.create_table(
"generated_videos",
@@ -180,9 +230,24 @@ def upgrade() -> None:
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.PrimaryKeyConstraint("id"),
)
op.create_index(op.f("ix_generated_videos_generation_task_id"), "generated_videos", ["generation_task_id"], unique=False)
op.create_index(op.f("ix_generated_videos_project_id"), "generated_videos", ["project_id"], unique=False)
op.create_index(op.f("ix_generated_videos_workspace_id"), "generated_videos", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_generated_videos_generation_task_id"),
"generated_videos",
["generation_task_id"],
unique=False,
)
op.create_index(
op.f("ix_generated_videos_project_id"),
"generated_videos",
["project_id"],
unique=False,
)
op.create_index(
op.f("ix_generated_videos_workspace_id"),
"generated_videos",
["workspace_id"],
unique=False,
)
op.create_table(
"tasks",
@@ -244,7 +309,12 @@ def upgrade() -> None:
)
op.create_index(op.f("ix_task_issues_project_id"), "task_issues", ["project_id"], unique=False)
op.create_index(op.f("ix_task_issues_task_id"), "task_issues", ["task_id"], unique=False)
op.create_index(op.f("ix_task_issues_workspace_id"), "task_issues", ["workspace_id"], unique=False)
op.create_index(
op.f("ix_task_issues_workspace_id"),
"task_issues",
["workspace_id"],
unique=False,
)
def downgrade() -> None:
+1 -2
View File
@@ -1,5 +1,3 @@
from fastapi import APIRouter
from app.api.routes.asset_libraries import router as asset_libraries_router
from app.api.routes.assets import router as assets_router
from app.api.routes.auth_simple import router as auth_router
@@ -11,6 +9,7 @@ from app.api.routes.ingest_jobs import router as ingest_jobs_router
from app.api.routes.project_management import router as project_management_router
from app.api.routes.projects import router as projects_router
from app.api.routes.upload import router as upload_router
from fastapi import APIRouter
api_router = APIRouter(prefix="/api/v1")
health_router = APIRouter()
+13 -4
View File
@@ -1,9 +1,18 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.dependencies import get_asset_library_repository
from app.schemas.asset_library import AssetLibraryResponse, CreateAssetLibraryRequest, ListAssetLibrariesResponse
from typing import Any
from packages.application import CreateAssetLibraryCommand, CreateAssetLibraryUseCase, ListAssetLibrariesUseCase
from app.schemas.asset_library import (
AssetLibraryResponse,
CreateAssetLibraryRequest,
ListAssetLibrariesResponse,
)
from fastapi import APIRouter, Depends
from packages.application import (
CreateAssetLibraryCommand,
CreateAssetLibraryUseCase,
ListAssetLibrariesUseCase,
)
from packages.domain import AssetLibraryKind
router = APIRouter()
+8 -3
View File
@@ -1,9 +1,14 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.dependencies import get_asset_repository
from app.schemas.asset import AssetResponse, CreateAssetRequest, ListAssetsResponse
from typing import Any
from packages.application import CreateAssetCommand, CreateAssetUseCase, ListAssetsUseCase
from fastapi import APIRouter, Depends
from packages.application import (
CreateAssetCommand,
CreateAssetUseCase,
ListAssetsUseCase,
)
from packages.domain import AssetStatus, ClassificationStatus
router = APIRouter()
+24 -20
View File
@@ -5,36 +5,35 @@ This module is intentionally disabled until the DI container/auth ports are rebu
Do not mount it directly; use `auth_simple.py` only as the current compatibility route.
"""
raise RuntimeError(
"apps.api.app.api.routes.auth is disabled: rebuild DI container before mounting full auth routes"
)
raise RuntimeError("apps.api.app.api.routes.auth is disabled: rebuild DI container before mounting full auth routes")
from fastapi import APIRouter, HTTPException, status, Depends, Request
from fastapi import APIRouter, Depends, HTTPException, Request, status
from pydantic import BaseModel, EmailStr
from packages.application.auth import (
RegisterUserUseCase,
RegisterUserRequest,
LoginUseCase,
LoginRequest,
LogoutUseCase,
LogoutRequest,
VerifyEmailUseCase,
VerifyEmailRequest,
RequestPasswordResetUseCase,
RequestPasswordResetRequest,
ResetPasswordUseCase,
ResetPasswordRequest,
)
from packages.domain.entities import User
from apps.api.app.dependencies import get_container
from apps.api.app.middleware.auth import get_current_user
from packages.application.auth import (
LoginRequest,
LoginUseCase,
LogoutRequest,
LogoutUseCase,
RegisterUserRequest,
RegisterUserUseCase,
RequestPasswordResetRequest,
RequestPasswordResetUseCase,
ResetPasswordRequest,
ResetPasswordUseCase,
VerifyEmailRequest,
VerifyEmailUseCase,
)
from packages.domain.entities import User
router = APIRouter(prefix="/auth", tags=["Authentication"])
# ==================== Request/Response Models ====================
class RegisterRequestModel(BaseModel):
email: EmailStr
password: str
@@ -77,7 +76,12 @@ 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 -4
View File
@@ -1,17 +1,18 @@
"""
认证 APISQLAlchemy ORM
"""
from datetime import datetime, timedelta, timezone
import hashlib
import secrets
from datetime import datetime, timedelta, timezone
import jwt
from app.config import settings
from app.dependencies import get_db_session
from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel, EmailStr
from sqlalchemy.orm import Session
from app.config import settings
from app.dependencies import get_db_session
from packages.adapters.sqlalchemy_impl.models import UserModel
from packages.domain.auth import password_hasher, password_validator
@@ -169,4 +170,7 @@ async def login(request: LoginRequest, db: Session = Depends(get_db_session)):
@router.get("/me")
async def get_current_user_info():
raise HTTPException(status_code=status.HTTP_501_NOT_IMPLEMENTED, detail="/auth/me requires bearer-token dependency integration")
raise HTTPException(
status_code=status.HTTP_501_NOT_IMPLEMENTED,
detail="/auth/me requires bearer-token dependency integration",
)
+11 -5
View File
@@ -1,12 +1,18 @@
from datetime import datetime, timezone
from fastapi import APIRouter, Depends, HTTPException
from typing import Any
from app.core.celery_app import celery_app
from app.dependencies import get_classification_job_repository
from app.schemas.classification_job import ClassificationJobResponse, SubmitClassificationJobRequest
from typing import Any
from packages.application import SubmitClassificationJobCommand, SubmitClassificationJobUseCase
from app.schemas.classification_job import (
ClassificationJobResponse,
SubmitClassificationJobRequest,
)
from fastapi import APIRouter, Depends, HTTPException
from packages.application import (
SubmitClassificationJobCommand,
SubmitClassificationJobUseCase,
)
router = APIRouter()
+3 -2
View File
@@ -1,4 +1,4 @@
from fastapi import APIRouter, Depends, HTTPException
from typing import Any
from app.core.storage import MinIOService, get_minio_service
from app.dependencies import get_generated_video_repository
@@ -7,7 +7,8 @@ from app.schemas.generated_video import (
GeneratedVideoResponse,
ListGeneratedVideosResponse,
)
from typing import Any
from fastapi import APIRouter, Depends, HTTPException
from packages.application import (
GetGeneratedVideoDownloadUrlUseCase,
GetGeneratedVideoUseCase,
+15 -5
View File
@@ -1,10 +1,20 @@
from fastapi import APIRouter, Depends, HTTPException
from typing import Any
from app.core.celery_app import celery_app
from app.dependencies import get_generation_task_repository, get_generated_video_repository
from app.schemas.generation_task import CreateGenerationTaskRequest, GenerationTaskResponse
from app.schemas.generated_video import GeneratedVideoResponse, ListGeneratedVideosResponse
from typing import Any
from app.dependencies import (
get_generated_video_repository,
get_generation_task_repository,
)
from app.schemas.generated_video import (
GeneratedVideoResponse,
ListGeneratedVideosResponse,
)
from app.schemas.generation_task import (
CreateGenerationTaskRequest,
GenerationTaskResponse,
)
from fastapi import APIRouter, Depends, HTTPException
from packages.application import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
+33 -13
View File
@@ -1,12 +1,11 @@
from pydantic import BaseModel
from datetime import datetime
import psycopg2
import redis
from app.config import settings
from fastapi import APIRouter, status
from fastapi.responses import JSONResponse
from app.config import settings
from pydantic import BaseModel
router = APIRouter(tags=["Health"])
@@ -56,16 +55,28 @@ async def startup_check():
async def _check_database() -> dict:
if settings.USE_IN_MEMORY_DB:
return {"status": "healthy", "type": "in_memory", "message": "Using in-memory database"}
return {
"status": "healthy",
"type": "in_memory",
"message": "Using in-memory database",
}
try:
conn = psycopg2.connect(settings.DATABASE_URL, connect_timeout=3)
with conn.cursor() as cur:
cur.execute("SELECT 1")
cur.fetchone()
conn.close()
return {"status": "healthy", "type": "postgresql", "message": "Database connection successful"}
return {
"status": "healthy",
"type": "postgresql",
"message": "Database connection successful",
}
except Exception as error:
return {"status": "unhealthy", "type": "postgresql", "message": f"Database connection failed: {error}"}
return {
"status": "unhealthy",
"type": "postgresql",
"message": f"Database connection failed: {error}",
}
async def _check_redis() -> dict:
@@ -73,23 +84,32 @@ async def _check_redis() -> dict:
client = redis.from_url(settings.REDIS_URL, socket_connect_timeout=3)
client.ping()
client.close()
return {"status": "healthy", "type": "redis", "message": "Redis connection successful"}
return {
"status": "healthy",
"type": "redis",
"message": "Redis connection successful",
}
except Exception as error:
return {"status": "unhealthy", "type": "redis", "message": f"Redis connection failed: {error}"}
return {
"status": "unhealthy",
"type": "redis",
"message": f"Redis connection failed: {error}",
}
async def _check_migrations() -> dict:
if settings.USE_IN_MEMORY_DB:
return {"status": "healthy", "message": "Using in-memory database, no migrations needed"}
return {
"status": "healthy",
"message": "Using in-memory database, no migrations needed",
}
try:
conn = psycopg2.connect(settings.DATABASE_URL, connect_timeout=3)
with conn.cursor() as cur:
cur.execute(
"""
cur.execute("""
SELECT COUNT(*) FROM information_schema.tables
WHERE table_name IN ('projects', 'asset_libraries', 'assets', 'ingest_jobs', 'classification_jobs')
"""
)
""")
count = cur.fetchone()[0]
conn.close()
if count >= 5:
+3 -2
View File
@@ -1,9 +1,10 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.core.celery_app import celery_app
from app.dependencies import get_ingest_job_repository
from app.schemas.ingest_job import IngestJobResponse, SubmitIngestJobRequest
from typing import Any
from fastapi import APIRouter, Depends
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
router = APIRouter()
@@ -1,4 +1,5 @@
"""项目管理 API 路由"""
from datetime import datetime
from typing import Annotated
@@ -11,7 +12,6 @@ from packages.adapters.sqlite_tracker.project_management_repositories import (
SQLiteTaskRepository,
)
from packages.application.get_task_detail_use_case import GetTaskDetailUseCase
from packages.application.update_task_use_case import UpdateTaskUseCase
from packages.application.project_management_use_cases import (
CreateMilestoneUseCase,
CreateTaskIssueUseCase,
@@ -23,6 +23,7 @@ from packages.application.project_management_use_cases import (
UpdateTaskProgressUseCase,
UpdateTaskStatusUseCase,
)
from packages.application.update_task_use_case import UpdateTaskUseCase
from packages.domain import TaskPriority, TaskStatus
router = APIRouter()
+13 -4
View File
@@ -1,9 +1,18 @@
from fastapi import APIRouter, Depends
from typing import Any
from app.dependencies import get_project_repository
from app.schemas.project import CreateProjectRequest, ListProjectsResponse, ProjectResponse
from typing import Any
from packages.application import CreateProjectCommand, CreateProjectUseCase, ListProjectsUseCase
from app.schemas.project import (
CreateProjectRequest,
ListProjectsResponse,
ProjectResponse,
)
from fastapi import APIRouter, Depends
from packages.application import (
CreateProjectCommand,
CreateProjectUseCase,
ListProjectsUseCase,
)
router = APIRouter()
+3 -2
View File
@@ -1,11 +1,12 @@
from fastapi import APIRouter, Depends, File, Form, UploadFile
from typing import Any
from uuid import uuid4
from app.core.celery_app import celery_app
from app.core.storage import MinIOService, get_minio_service
from app.dependencies import get_ingest_job_repository
from app.schemas.upload import UploadAssetResponse
from typing import Any
from fastapi import APIRouter, Depends, File, Form, UploadFile
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
router = APIRouter()
+16 -5
View File
@@ -9,21 +9,28 @@ raise RuntimeError(
"apps.api.app.api.routes.workspaces is disabled: rebuild DI container before mounting workspace routes"
)
from fastapi import APIRouter, HTTPException, status, Depends
from pydantic import BaseModel, EmailStr
from typing import List
from datetime import datetime
from typing import List
from fastapi import APIRouter, Depends, HTTPException, status
from pydantic import BaseModel, EmailStr
from apps.api.app.dependencies import get_container
from apps.api.app.middleware.auth import (
get_current_user,
require_workspace_access,
require_workspace_admin,
require_workspace_owner,
)
from packages.application.workspace import *
from packages.domain.entities import User
from apps.api.app.dependencies import get_container
from apps.api.app.middleware.auth import get_current_user, require_workspace_access, require_workspace_admin, require_workspace_owner
router = APIRouter(prefix="/workspaces", tags=["Workspaces"])
# ==================== Request/Response Models ====================
class CreateWorkspaceRequestModel(BaseModel):
name: str
subscription_plan: str = "free"
@@ -52,6 +59,7 @@ class UpgradeSubscriptionRequestModel(BaseModel):
# ==================== Workspace CRUD ====================
@router.post("", response_model=WorkspaceResponseModel, status_code=status.HTTP_201_CREATED)
async def create_workspace(
request: CreateWorkspaceRequestModel,
@@ -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,
@@ -272,6 +281,7 @@ async def update_member_role(
# ==================== Subscription Management ====================
@router.post("/{workspace_id}/subscription/upgrade")
async def upgrade_subscription(
workspace_id: str,
@@ -346,6 +356,7 @@ async def get_quota_status(
# ==================== Invitation Acceptance ====================
@router.post("/invitations/{token}/accept")
async def accept_invitation(
token: str,
+3 -2
View File
@@ -1,6 +1,7 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
from typing import Optional
import os
from typing import Optional
from pydantic_settings import BaseSettings, SettingsConfigDict
class Settings(BaseSettings):
+1 -3
View File
@@ -1,7 +1,5 @@
from celery import Celery
from app.config import get_settings
from celery import Celery
settings = get_settings()
celery_app = Celery("xiaoxia-saas-api")
+3 -2
View File
@@ -1,6 +1,7 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
from typing import Optional
import os
from typing import Optional
from pydantic_settings import BaseSettings, SettingsConfigDict
class DatabaseSettings(BaseSettings):
+9 -13
View File
@@ -1,7 +1,8 @@
"""阿里云 OSS 存储服务"""
import logging
from urllib.parse import urlparse
import os
from urllib.parse import urlparse
try:
import oss2
@@ -63,19 +64,11 @@ 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:
@@ -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)
@@ -145,7 +138,10 @@ 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:
"""
+6 -2
View File
@@ -1,9 +1,13 @@
from collections.abc import Generator
from app.config import settings
from sqlalchemy.orm import Session
from app.config import settings
from packages.adapters.sqlalchemy_impl import build_session_factory, ensure_database_exists, initialize_database
from packages.adapters.sqlalchemy_impl import (
build_session_factory,
ensure_database_exists,
initialize_database,
)
ensure_database_exists(settings.DATABASE_URL)
engine, SessionLocal = build_session_factory(
+40 -14
View File
@@ -1,14 +1,26 @@
from app.config import settings
from fastapi import Depends
from sqlalchemy.orm import Session
from app.config import settings
from packages.adapters.sqlalchemy_impl.asset_library_repository import SQLAlchemyAssetLibraryRepository
from packages.adapters.sqlalchemy_impl.asset_library_repository import (
SQLAlchemyAssetLibraryRepository,
)
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from packages.adapters.sqlalchemy_impl.classification_job_repository import SQLAlchemyClassificationJobRepository
from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository
from packages.adapters.sqlalchemy_impl.generation_task_repository import SQLAlchemyGenerationTaskRepository
from packages.adapters.sqlalchemy_impl.ingest_job_repository import SQLAlchemyIngestJobRepository
from packages.adapters.sqlalchemy_impl.project_repository import SQLAlchemyProjectRepository
from packages.adapters.sqlalchemy_impl.classification_job_repository import (
SQLAlchemyClassificationJobRepository,
)
from packages.adapters.sqlalchemy_impl.generated_video_repository import (
SQLAlchemyGeneratedVideoRepository,
)
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
from packages.adapters.sqlalchemy_impl.ingest_job_repository import (
SQLAlchemyIngestJobRepository,
)
from packages.adapters.sqlalchemy_impl.project_repository import (
SQLAlchemyProjectRepository,
)
from packages.adapters.sqlalchemy_impl.session import build_session_factory
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
@@ -22,29 +34,43 @@ def get_db_session():
session.close()
def get_asset_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyAssetRepository:
def get_asset_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyAssetRepository:
return SQLAlchemyAssetRepository(session)
def get_asset_library_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyAssetLibraryRepository:
def get_asset_library_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyAssetLibraryRepository:
return SQLAlchemyAssetLibraryRepository(session)
def get_ingest_job_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyIngestJobRepository:
def get_ingest_job_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyIngestJobRepository:
return SQLAlchemyIngestJobRepository(session)
def get_classification_job_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyClassificationJobRepository:
def get_classification_job_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyClassificationJobRepository:
return SQLAlchemyClassificationJobRepository(session)
def get_generation_task_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyGenerationTaskRepository:
def get_generation_task_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyGenerationTaskRepository:
return SQLAlchemyGenerationTaskRepository(session)
def get_generated_video_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyGeneratedVideoRepository:
def get_generated_video_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyGeneratedVideoRepository:
return SQLAlchemyGeneratedVideoRepository(session)
def get_project_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyProjectRepository:
def get_project_repository(
session: Session = Depends(get_db_session),
) -> SQLAlchemyProjectRepository:
return SQLAlchemyProjectRepository(session)
+3 -6
View File
@@ -5,17 +5,14 @@ Disabled because it depends on the removed DI container. Rebuild it around the
canonical JWT settings and SQLAlchemy-backed user repository before reuse.
"""
raise RuntimeError(
"apps.api.app.middleware.auth is disabled: rebuild auth dependency wiring before importing it"
)
raise RuntimeError("apps.api.app.middleware.auth is disabled: rebuild auth dependency wiring before importing it")
from fastapi import Depends, HTTPException, status
from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from apps.api.app.dependencies import get_container
from packages.domain.auth import jwt_service
from packages.domain.entities import User
from apps.api.app.dependencies import get_container
security = HTTPBearer()
+11 -7
View File
@@ -1,12 +1,14 @@
"""
全局异常处理和错误响应
"""
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__)
@@ -100,11 +102,13 @@ async def validation_exception_handler(request: Request, exc: RequestValidationE
"""请求验证异常处理"""
errors = []
for error in exc.errors():
errors.append({
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,
+7 -8
View File
@@ -1,8 +1,10 @@
"""
请求日志中间件
"""
import time
import logging
import time
from fastapi import Request
from starlette.middleware.base import BaseHTTPMiddleware
@@ -27,8 +29,7 @@ class RequestLoggingMiddleware(BaseHTTPMiddleware):
# 记录响应信息
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"
)
# 添加响应头
@@ -54,8 +55,7 @@ class RateLimitMiddleware(BaseHTTPMiddleware):
# 清理过期记录
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
]
# 检查速率限制
@@ -63,6 +63,7 @@ class RateLimitMiddleware(BaseHTTPMiddleware):
if request_count >= self.max_requests:
from fastapi.responses import JSONResponse
return JSONResponse(
status_code=429,
content={
@@ -83,8 +84,6 @@ class RateLimitMiddleware(BaseHTTPMiddleware):
# 添加速率限制信息到响应头
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
+10 -8
View File
@@ -1,9 +1,11 @@
"""
性能监控中间件
"""
import time
import logging
import time
from typing import Callable
from fastapi import Request, Response
from starlette.middleware.base import BaseHTTPMiddleware
@@ -59,13 +61,14 @@ class PerformanceMonitoringMiddleware(BaseHTTPMiddleware):
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())
@@ -78,19 +81,18 @@ class DatabaseQueryLogger:
def log_query(self, query: str, params: tuple, duration: float):
"""记录查询"""
self.queries.append({
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):
"""获取统计信息"""
+7 -6
View File
@@ -1,9 +1,11 @@
"""
API 版本管理中间件
"""
from datetime import datetime
from fastapi import Request
from starlette.middleware.base import BaseHTTPMiddleware
from datetime import datetime
class APIVersionMiddleware(BaseHTTPMiddleware):
@@ -45,9 +47,7 @@ class APIVersionMiddleware(BaseHTTPMiddleware):
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
@@ -70,6 +70,7 @@ class VersionNotFoundMiddleware(BaseHTTPMiddleware):
if version in self.SUNSET_VERSIONS:
from fastapi.responses import JSONResponse
return JSONResponse(
status_code=410,
content={
@@ -77,9 +78,9 @@ 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)
+5 -1
View File
@@ -1,7 +1,11 @@
"""Schema package."""
from .asset import AssetResponse, CreateAssetRequest, ListAssetsResponse
from .asset_library import AssetLibraryResponse, CreateAssetLibraryRequest, ListAssetLibrariesResponse
from .asset_library import (
AssetLibraryResponse,
CreateAssetLibraryRequest,
ListAssetLibrariesResponse,
)
from .health import HealthResponse
from .ingest_job import IngestJobResponse, SubmitIngestJobRequest
from .project import CreateProjectRequest, ListProjectsResponse, ProjectResponse
+6 -7
View File
@@ -1,12 +1,5 @@
import os
from fastapi import FastAPI
from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from fastapi.middleware.gzip import GZipMiddleware
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.staticfiles import StaticFiles
from app.api.router import api_router, health_router
from app.config import settings
from app.middleware.exceptions import (
@@ -17,6 +10,12 @@ from app.middleware.exceptions import (
validation_exception_handler,
)
from app.middleware.logging import RequestLoggingMiddleware
from fastapi import FastAPI
from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from fastapi.middleware.gzip import GZipMiddleware
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.staticfiles import StaticFiles
app = FastAPI(
title="小虾 SaaS API",
+15 -5
View File
@@ -4,13 +4,22 @@ from datetime import datetime, timezone
from app.config import get_settings
from app.core.storage import get_minio_service
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from packages.adapters.sqlalchemy_impl.generated_video_repository import (
SQLAlchemyGeneratedVideoRepository,
)
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
from packages.adapters.sqlalchemy_impl.session import (
SessionLocal,
build_session_factory,
)
from packages.domain import GeneratedVideo, GenerationTaskStatus
from .celery_app import celery_app
from .video_processing import VideoProcessor
from packages.adapters.sqlalchemy_impl.session import SessionLocal, build_session_factory
from packages.adapters.sqlalchemy_impl.generation_task_repository import SQLAlchemyGenerationTaskRepository
from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from packages.domain import GeneratedVideo, GenerationTaskStatus
settings = get_settings()
if SessionLocal is None:
@@ -151,6 +160,7 @@ def generate_video(task_id: str) -> dict:
# 清理临时文件
try:
import shutil
shutil.rmtree(temp_dir, ignore_errors=True)
except:
pass
+1
View File
@@ -1,6 +1,7 @@
"""
视频处理模块
"""
from .processor import VideoProcessor, VideoResult
__all__ = ["VideoProcessor", "VideoResult"]
+4 -4
View File
@@ -1,6 +1,7 @@
"""
视频处理核心类
"""
import os
import tempfile
from dataclasses import dataclass
@@ -13,6 +14,7 @@ import ffmpeg
@dataclass
class VideoResult:
"""视频生成结果"""
output_path: str
thumbnail_path: str
duration: float
@@ -70,8 +72,7 @@ class VideoProcessor:
# 使用 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",
@@ -142,8 +143,7 @@ class VideoProcessor:
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)
-2
View File
@@ -1,8 +1,6 @@
from celery import Celery
from worker_app.core.config import get_settings
settings = get_settings()
celery_app = Celery(settings.worker_name)
celery_app.conf.broker_url = settings.broker_url
+3 -2
View File
@@ -1,6 +1,7 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
from typing import Optional
import os
from typing import Optional
from pydantic_settings import BaseSettings, SettingsConfigDict
class WorkerSettings(BaseSettings):
+6 -1
View File
@@ -1,5 +1,10 @@
from worker_app.core.config import get_settings
from packages.adapters.sqlalchemy_impl import build_session_factory, ensure_database_exists, initialize_database
from packages.adapters.sqlalchemy_impl import (
build_session_factory,
ensure_database_exists,
initialize_database,
)
settings = get_settings()
ensure_database_exists(settings.database_url)
@@ -1,7 +1,14 @@
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")
+19 -6
View File
@@ -7,6 +7,8 @@ from pathlib import Path
from urllib.parse import urlparse
import oss2
from worker_app.celery_app import celery_app
from worker_app.db import SessionLocal
from packages.adapters.sqlalchemy_impl import (
SQLAlchemyAssetRepository,
@@ -14,8 +16,6 @@ from packages.adapters.sqlalchemy_impl import (
SQLAlchemyGenerationTaskRepository,
)
from packages.domain import GeneratedVideo, GenerationTaskStatus
from worker_app.celery_app import celery_app
from worker_app.db import SessionLocal
OUTPUT_WIDTH = 1280
OUTPUT_HEIGHT = 720
@@ -176,7 +176,11 @@ def generate_video(task_id: str) -> dict:
task = task_repo.get(task_id)
if task is None:
db.close()
return {"status": "failed", "error": "generation task not found", "task_id": task_id}
return {
"status": "failed",
"error": "generation task not found",
"task_id": task_id,
}
try:
task.status = GenerationTaskStatus.RUNNING
@@ -184,10 +188,14 @@ def generate_video(task_id: str) -> dict:
task.started_at = task.started_at or datetime.now(timezone.utc)
task_repo.update(task)
assets = [asset for asset in asset_repo.list_by_library(task.asset_library_id) if asset.mime_type.startswith("video")]
assets = [
asset for asset in asset_repo.list_by_library(task.asset_library_id) if asset.mime_type.startswith("video")
]
output_name = f"generated-{task.id}.mp4"
storage_key = f"generated/workspaces/{task.workspace_id}/projects/{task.project_id}/tasks/{task.id}/{output_name}"
storage_key = (
f"generated/workspaces/{task.workspace_id}/projects/{task.project_id}/tasks/{task.id}/{output_name}"
)
with tempfile.TemporaryDirectory(prefix="xiaoxia-generation-") as temp_dir:
temp_path = Path(temp_dir)
@@ -235,7 +243,12 @@ def generate_video(task_id: str) -> dict:
task.completed_at = datetime.now(timezone.utc)
task_repo.update(task)
return {"status": "completed", "task_id": task.id, "video_id": video.id, "file_url": file_url}
return {
"status": "completed",
"task_id": task.id,
"video_id": video.id,
"file_url": file_url,
}
except Exception as error:
task.status = GenerationTaskStatus.FAILED
task.error_message = str(error)
+6 -2
View File
@@ -1,10 +1,14 @@
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:
@@ -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,7 +1,9 @@
"""
用户仓储 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
@@ -1,8 +1,10 @@
"""
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
@@ -1,7 +1,9 @@
"""
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
@@ -1,7 +1,9 @@
"""
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
+1 -3
View File
@@ -5,6 +5,4 @@ package is kept only as a migration marker; do not import it in application,
API, worker, or new tests.
"""
raise RuntimeError(
"packages.adapters.postgres is deprecated; use packages.adapters.sqlalchemy_impl instead"
)
raise RuntimeError("packages.adapters.postgres is deprecated; use packages.adapters.sqlalchemy_impl instead")
+19 -5
View File
@@ -1,6 +1,7 @@
"""
Asset PostgreSQL Repository 实现
"""
import json
from sqlalchemy import and_, func, select
@@ -24,7 +25,7 @@ class PostgresAssetRepository(AssetRepository):
project_id=asset.project_id,
asset_library_id=asset.library_id,
name=asset.name,
file_type=asset.mime_type.split("/")[0] if "/" in asset.mime_type else asset.mime_type,
file_type=(asset.mime_type.split("/")[0] if "/" in asset.mime_type else asset.mime_type),
file_size=asset.file_size,
file_url=asset.storage_key,
thumbnail_url=asset.thumbnail_url,
@@ -35,7 +36,7 @@ class PostgresAssetRepository(AssetRepository):
codec=asset.codec,
status=asset.status.value,
classification_status=asset.classification_status.value,
classification_result=json.dumps(asset.metadata) if asset.metadata else None,
classification_result=(json.dumps(asset.metadata) if asset.metadata else None),
quality_score=asset.quality_score,
uploaded_by_user_id=asset.uploaded_by_user_id or "system",
created_at=asset.created_at,
@@ -59,7 +60,12 @@ class PostgresAssetRepository(AssetRepository):
) -> list[Asset]:
result = await self.session.execute(
select(AssetModel)
.where(and_(AssetModel.project_id == project_id, AssetModel.workspace_id == workspace_id))
.where(
and_(
AssetModel.project_id == project_id,
AssetModel.workspace_id == workspace_id,
)
)
.order_by(AssetModel.created_at.desc())
.offset(skip)
.limit(limit)
@@ -75,7 +81,12 @@ class PostgresAssetRepository(AssetRepository):
) -> list[Asset]:
result = await self.session.execute(
select(AssetModel)
.where(and_(AssetModel.asset_library_id == library_id, AssetModel.workspace_id == workspace_id))
.where(
and_(
AssetModel.asset_library_id == library_id,
AssetModel.workspace_id == workspace_id,
)
)
.order_by(AssetModel.created_at.desc())
.offset(skip)
.limit(limit)
@@ -118,7 +129,10 @@ class PostgresAssetRepository(AssetRepository):
async def count_by_project(self, project_id: str, workspace_id: str) -> int:
result = await self.session.execute(
select(func.count(AssetModel.id)).where(
and_(AssetModel.project_id == project_id, AssetModel.workspace_id == workspace_id)
and_(
AssetModel.project_id == project_id,
AssetModel.workspace_id == workspace_id,
)
)
)
return result.scalar() or 0
@@ -1,7 +1,9 @@
"""
数据库连接池管理
"""
from typing import Optional
import psycopg2
from psycopg2 import pool
from psycopg2.extras import RealDictCursor
@@ -10,7 +12,7 @@ from psycopg2.extras import RealDictCursor
class DatabaseConnectionPool:
"""PostgreSQL 连接池"""
_instance: Optional['DatabaseConnectionPool'] = None
_instance: Optional["DatabaseConnectionPool"] = None
_pool: Optional[pool.ThreadedConnectionPool] = None
def __new__(cls):
@@ -1,7 +1,9 @@
"""
PostgreSQL Project Repository 实现
"""
from typing import Optional, List
from typing import List, Optional
import psycopg2
from psycopg2.extras import RealDictCursor
@@ -18,6 +20,7 @@ class PostgresProjectRepository(ProjectRepository):
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, project: Project) -> None:
@@ -25,7 +28,8 @@ class PostgresProjectRepository(ProjectRepository):
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,7 +42,8 @@ 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,
@@ -47,7 +52,8 @@ class PostgresProjectRepository(ProjectRepository):
"created_by": project.created_by,
"created_at": project.created_at,
"updated_at": project.updated_at,
})
},
)
conn.commit()
finally:
conn.close()
@@ -70,7 +76,7 @@ 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]
@@ -84,7 +90,7 @@ 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]
@@ -98,7 +104,7 @@ 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:
+10 -4
View File
@@ -1,10 +1,12 @@
"""
PostgreSQL User Repository 实现
"""
from datetime import datetime
from typing import Optional
import psycopg2
from psycopg2.extras import RealDictCursor
from datetime import datetime
from packages.domain.entities import User
from packages.ports.user_repository import UserRepository
@@ -19,6 +21,7 @@ class PostgresUserRepository(UserRepository):
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, user: User) -> None:
@@ -27,7 +30,8 @@ class PostgresUserRepository(UserRepository):
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,7 +54,8 @@ 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,
@@ -63,7 +68,8 @@ class PostgresUserRepository(UserRepository):
"last_login_at": user.last_login_at,
"last_login_ip": user.last_login_ip,
"created_at": user.created_at,
})
},
)
conn.commit()
finally:
conn.close()
@@ -1,7 +1,9 @@
"""
PostgreSQL WorkspaceInvitation Repository 实现
"""
from typing import Optional, List
from typing import List, Optional
import psycopg2
from psycopg2.extras import RealDictCursor
@@ -18,6 +20,7 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, invitation: WorkspaceInvitation) -> None:
@@ -25,7 +28,8 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
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,7 +39,8 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
)
ON CONFLICT (id) DO UPDATE SET
status = EXCLUDED.status
""", {
""",
{
"id": invitation.id,
"workspace_id": invitation.workspace_id,
"email": invitation.email,
@@ -45,7 +50,8 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
"expires_at": invitation.expires_at,
"status": invitation.status,
"created_at": invitation.created_at,
})
},
)
conn.commit()
finally:
conn.close()
@@ -55,7 +61,10 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
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:
@@ -79,7 +88,7 @@ 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]
@@ -91,11 +100,14 @@ class PostgresWorkspaceInvitationRepository(WorkspaceInvitationRepository):
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:
@@ -1,7 +1,9 @@
"""
PostgreSQL WorkspaceMember Repository 实现
"""
from typing import Optional, List
from typing import List, Optional
import psycopg2
from psycopg2.extras import RealDictCursor
@@ -18,6 +20,7 @@ class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, member: WorkspaceMember) -> None:
@@ -25,7 +28,8 @@ class PostgresWorkspaceMemberRepository(WorkspaceMemberRepository):
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,14 +38,16 @@ 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,
})
},
)
conn.commit()
finally:
conn.close()
@@ -68,7 +74,7 @@ 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
@@ -82,7 +88,7 @@ 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]
@@ -96,7 +102,7 @@ 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]
@@ -110,7 +116,7 @@ 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:
@@ -1,7 +1,9 @@
"""
PostgreSQL Workspace Repository 实现
"""
from typing import Optional
import psycopg2
from psycopg2.extras import RealDictCursor
@@ -18,6 +20,7 @@ class PostgresWorkspaceRepository(WorkspaceRepository):
def _get_connection(self):
"""获取数据库连接(使用连接池)"""
from packages.adapters.postgres.connection_pool import PooledConnection
return PooledConnection()
def save(self, workspace: Workspace) -> None:
@@ -25,7 +28,8 @@ class PostgresWorkspaceRepository(WorkspaceRepository):
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,7 +47,8 @@ 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,
@@ -54,7 +59,8 @@ class PostgresWorkspaceRepository(WorkspaceRepository):
"max_storage_gb": workspace.max_storage_gb,
"used_storage_gb": workspace.used_storage_gb,
"created_at": workspace.created_at,
})
},
)
conn.commit()
finally:
conn.close()
+5 -1
View File
@@ -1,3 +1,7 @@
from packages.adapters.redis.session_store import RedisConfig, SessionStore, get_session_store
from packages.adapters.redis.session_store import (
RedisConfig,
SessionStore,
get_session_store,
)
__all__ = ["RedisConfig", "SessionStore", "get_session_store"]
+8 -17
View File
@@ -2,15 +2,18 @@
Redis Session 存储
用于存储 refresh_token 和 Session 信息
"""
from typing import Optional
from datetime import datetime, timedelta, timezone
import json
from datetime import datetime, timedelta, timezone
from typing import Optional
import redis
from redis import Redis
class RedisConfig:
"""Redis 配置"""
HOST: str = "localhost"
PORT: int = 6379
DB: int = 0
@@ -92,19 +95,11 @@ class SessionStore:
# 保存 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)
@@ -175,11 +170,7 @@ class SessionStore:
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
+5 -1
View File
@@ -1,3 +1,7 @@
from packages.adapters.smtp.email_service import EmailConfig, EmailService, get_email_service
from packages.adapters.smtp.email_service import (
EmailConfig,
EmailService,
get_email_service,
)
__all__ = ["EmailConfig", "EmailService", "get_email_service"]
+6 -8
View File
@@ -2,16 +2,18 @@
邮件服务
支持 SMTP 发送邮件(验证/重置密码/邀请等)
"""
import smtplib
from email.mime.text import MIMEText
from email.mime.multipart import MIMEMultipart
from typing import Optional, List
from dataclasses import dataclass
from email.mime.multipart import MIMEMultipart
from email.mime.text import MIMEText
from typing import List, Optional
@dataclass
class EmailConfig:
"""邮件配置"""
smtp_host: str = "smtp.gmail.com"
smtp_port: int = 587
smtp_user: str = ""
@@ -91,11 +93,7 @@ class EmailService:
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
@@ -7,7 +7,13 @@ from .generated_video_repository import SQLAlchemyGeneratedVideoRepository
from .generation_task_repository import SQLAlchemyGenerationTaskRepository
from .ingest_job_repository import SQLAlchemyIngestJobRepository
from .project_repository import SQLAlchemyProjectRepository
from .session import Base, build_engine, build_session_factory, ensure_database_exists, initialize_database
from .session import (
Base,
build_engine,
build_session_factory,
ensure_database_exists,
initialize_database,
)
__all__ = [
"Base",
@@ -29,7 +29,7 @@ class SQLAlchemyAssetRepository:
project_id=asset.project_id,
asset_library_id=asset.library_id,
name=asset.name,
file_type=asset.mime_type.split('/')[0] if '/' in asset.mime_type else asset.mime_type,
file_type=(asset.mime_type.split("/")[0] if "/" in asset.mime_type else asset.mime_type),
file_size=asset.file_size,
file_url=asset.storage_key,
thumbnail_url=asset.thumbnail_url,
@@ -40,9 +40,9 @@ class SQLAlchemyAssetRepository:
codec=asset.codec,
status=asset.status.value,
classification_status=asset.classification_status.value,
classification_result=json.dumps(asset.metadata) if asset.metadata else None,
classification_result=(json.dumps(asset.metadata) if asset.metadata else None),
quality_score=asset.quality_score,
uploaded_by_user_id=asset.uploaded_by_user_id or 'system',
uploaded_by_user_id=asset.uploaded_by_user_id or "system",
created_at=asset.created_at,
updated_at=now,
)
@@ -80,11 +80,11 @@ class SQLAlchemyAssetRepository:
except Exception:
metadata = {}
mime_type = model.file_type
if '/' not in mime_type:
if "/" not in mime_type:
mime_type = {
'video': 'video/mp4',
'audio': 'audio/mpeg',
'image': 'image/jpeg',
"video": "video/mp4",
"audio": "audio/mpeg",
"image": "image/jpeg",
}.get(mime_type, mime_type)
return Asset(
id=model.id,
@@ -55,5 +55,9 @@ class SQLAlchemyGeneratedVideoRepository:
return [self.get(model.id) for model in models if self.get(model.id) is not None]
def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]:
models = self.session.query(GeneratedVideoModel).filter(GeneratedVideoModel.generation_task_id == generation_task_id).all()
models = (
self.session.query(GeneratedVideoModel)
.filter(GeneratedVideoModel.generation_task_id == generation_task_id)
.all()
)
return [self.get(model.id) for model in models if self.get(model.id) is not None]
+2 -1
View File
@@ -1,6 +1,7 @@
from datetime import datetime, timezone
from sqlalchemy import Boolean, Column, DateTime, Float, String, Text, create_engine
from sqlalchemy.orm import declarative_base
from datetime import datetime, timezone
Base = declarative_base()
@@ -1,5 +1,7 @@
"""项目管理 SQLAlchemy Repository 实现"""
import json
from sqlalchemy.orm import Session
from packages.domain import Milestone, Task, TaskIssue
@@ -8,6 +10,7 @@ from packages.ports.project_management_repositories import (
TaskIssueRepository,
TaskRepository,
)
from .models import MilestoneModel, TaskIssueModel, TaskModel
@@ -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,
+4 -2
View File
@@ -6,7 +6,6 @@ from sqlalchemy.orm import sessionmaker
from packages.adapters.sqlalchemy_impl.models import Base
SCHEMA_INIT_LOCK_ID = 2026061501
SessionLocal = None
@@ -77,5 +76,8 @@ def initialize_database(engine) -> None:
Base.metadata.create_all(bind=connection)
connection.commit()
finally:
connection.execute(text("SELECT pg_advisory_unlock(:lock_id)"), {"lock_id": SCHEMA_INIT_LOCK_ID})
connection.execute(
text("SELECT pg_advisory_unlock(:lock_id)"),
{"lock_id": SCHEMA_INIT_LOCK_ID},
)
connection.commit()
+2 -1
View File
@@ -1,8 +1,9 @@
"""SQLite Tracker Adapter"""
from .project_management_repositories import (
SQLiteTaskRepository,
SQLiteMilestoneRepository,
SQLiteTaskIssueRepository,
SQLiteTaskRepository,
)
__all__ = [
@@ -1,11 +1,20 @@
"""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 的任务仓储"""
@@ -22,17 +31,17 @@ class SQLiteTaskRepository:
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]:
@@ -41,30 +50,35 @@ class SQLiteTaskRepository:
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,
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",
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(),
)
)
return tasks
@@ -74,19 +88,38 @@ class SQLiteTaskRepository:
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()
@@ -108,15 +141,17 @@ class SQLiteMilestoneRepository:
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",
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()
))
created_at=(datetime.fromisoformat(row["start_date"]) if row["start_date"] else datetime.now()),
)
)
return milestones
@@ -125,17 +160,33 @@ class SQLiteMilestoneRepository:
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()
+14 -3
View File
@@ -1,15 +1,26 @@
"""Application use cases package."""
from .asset_libraries import CreateAssetLibraryCommand, CreateAssetLibraryUseCase, ListAssetLibrariesUseCase
from .asset_libraries import (
CreateAssetLibraryCommand,
CreateAssetLibraryUseCase,
ListAssetLibrariesUseCase,
)
from .assets import CreateAssetCommand, CreateAssetUseCase, ListAssetsUseCase
from .classification_jobs import SubmitClassificationJobCommand, SubmitClassificationJobUseCase
from .classification_jobs import (
SubmitClassificationJobCommand,
SubmitClassificationJobUseCase,
)
from .generated_videos import (
GetGeneratedVideoDownloadUrlUseCase,
GetGeneratedVideoUseCase,
ListGeneratedVideosByTaskUseCase,
ListGeneratedVideosUseCase,
)
from .generation_tasks import CreateGenerationTaskCommand, CreateGenerationTaskUseCase, GetGenerationTaskUseCase
from .generation_tasks import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
GetGenerationTaskUseCase,
)
from .ingest_jobs import SubmitIngestJobCommand, SubmitIngestJobUseCase
from .projects import CreateProjectCommand, CreateProjectUseCase, ListProjectsUseCase
+14 -13
View File
@@ -1,25 +1,26 @@
"""认证相关 Use Cases"""
from packages.application.auth.register_user_use_case import (
RegisterUserUseCase,
RegisterUserRequest,
RegisterUserResponse,
VerifyEmailUseCase,
VerifyEmailRequest,
)
from packages.application.auth.login_use_case import (
LoginUseCase,
LoginRequest,
LoginResponse,
RefreshTokenUseCase,
RefreshTokenRequest,
LogoutUseCase,
LoginUseCase,
LogoutRequest,
LogoutUseCase,
RefreshTokenRequest,
RefreshTokenUseCase,
)
from packages.application.auth.password_reset_use_case import (
RequestPasswordResetUseCase,
RequestPasswordResetRequest,
ResetPasswordUseCase,
RequestPasswordResetUseCase,
ResetPasswordRequest,
ResetPasswordUseCase,
)
from packages.application.auth.register_user_use_case import (
RegisterUserRequest,
RegisterUserResponse,
RegisterUserUseCase,
VerifyEmailRequest,
VerifyEmailUseCase,
)
__all__ = [
+10 -8
View File
@@ -1,15 +1,13 @@
"""
用户登录 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:
@@ -96,6 +94,7 @@ class LoginUseCase:
# 这里使用一个特殊的 "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,
@@ -124,7 +123,8 @@ class LoginUseCase:
self.user_repository.save(user)
# 9. 返回响应
return LoginResponse(
return (
LoginResponse(
access_token=access_token,
refresh_token=refresh_token,
user_id=user.id,
@@ -132,7 +132,9 @@ class LoginUseCase:
username=user.username,
display_name=user.display_name,
expires_in=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
), None
),
None,
)
except Exception as e:
return None, f"Login failed: {str(e)}"
@@ -1,6 +1,7 @@
"""
密码重置 Use Case
"""
import secrets
from datetime import datetime, timedelta, timezone
from typing import Optional
@@ -58,9 +59,7 @@ class RequestPasswordResetUseCase:
# 设置令牌和过期时间
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)
@@ -1,14 +1,15 @@
"""
用户注册 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:
@@ -140,13 +141,16 @@ class RegisterUserUseCase:
print(f"Email service error: {e}")
# 10. 返回响应(即使邮件发送失败,用户也已创建)
return RegisterUserResponse(
return (
RegisterUserResponse(
user_id=user.id,
email=user.email,
username=user.username,
display_name=user.display_name,
email_verification_sent=email_sent,
), None
),
None,
)
except Exception as e:
return None, f"Registration failed: {str(e)}"
+7 -3
View File
@@ -1,16 +1,18 @@
"""
通用分页器
"""
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")
@@ -27,6 +29,7 @@ class PaginationParams(BaseModel):
class PaginationMeta(BaseModel):
"""分页元数据"""
page: int = Field(..., description="当前页码")
page_size: int = Field(..., description="每页数量")
total: int = Field(..., description="总记录数")
@@ -55,6 +58,7 @@ class PaginationMeta(BaseModel):
class PaginatedResponse(BaseModel, Generic[T]):
"""分页响应"""
data: List[T] = Field(..., description="数据列表")
pagination: PaginationMeta = Field(..., description="分页信息")
@@ -1,4 +1,5 @@
"""获取单个任务详情用例"""
from packages.domain import Task
from packages.ports import TaskRepository
@@ -1,4 +1,5 @@
"""项目管理 Use Cases"""
from packages.domain import Milestone, Task, TaskIssue, TaskPriority, TaskStatus
from packages.ports import MilestoneRepository, TaskIssueRepository, TaskRepository
@@ -1,4 +1,5 @@
"""更新任务基本信息用例"""
from packages.domain import Task
from packages.ports import TaskRepository
+35 -34
View File
@@ -1,53 +1,54 @@
"""Workspace 相关 Use Cases"""
from packages.application.workspace.create_workspace_use_case import (
CreateWorkspaceUseCase,
CreateWorkspaceRequest,
CreateWorkspaceResponse,
)
from packages.application.workspace.invite_member_use_case import (
InviteMemberUseCase,
InviteMemberRequest,
InviteMemberResponse,
)
from packages.application.workspace.accept_invitation_use_case import (
AcceptInvitationUseCase,
AcceptInvitationRequest,
AcceptInvitationResponse,
DeclineInvitationUseCase,
AcceptInvitationUseCase,
DeclineInvitationRequest,
DeclineInvitationUseCase,
)
from packages.application.workspace.remove_member_use_case import (
RemoveMemberUseCase,
RemoveMemberRequest,
LeaveWorkspaceUseCase,
LeaveWorkspaceRequest,
from packages.application.workspace.create_workspace_use_case import (
CreateWorkspaceRequest,
CreateWorkspaceResponse,
CreateWorkspaceUseCase,
)
from packages.application.workspace.update_member_role_use_case import (
UpdateMemberRoleUseCase,
UpdateMemberRoleRequest,
UpdateMemberRoleResponse,
)
from packages.application.workspace.list_workspaces_use_case import (
ListWorkspacesUseCase,
ListWorkspacesRequest,
ListWorkspacesResponse,
GetWorkspaceDetailUseCase,
GetWorkspaceDetailRequest,
WorkspaceInfo,
WorkspaceDetailInfo,
from packages.application.workspace.invite_member_use_case import (
InviteMemberRequest,
InviteMemberResponse,
InviteMemberUseCase,
)
from packages.application.workspace.list_members_use_case import (
ListMembersUseCase,
ListMembersRequest,
ListMembersResponse,
ListMembersUseCase,
MemberInfo,
)
from packages.application.workspace.list_workspaces_use_case import (
GetWorkspaceDetailRequest,
GetWorkspaceDetailUseCase,
ListWorkspacesRequest,
ListWorkspacesResponse,
ListWorkspacesUseCase,
WorkspaceDetailInfo,
WorkspaceInfo,
)
from packages.application.workspace.remove_member_use_case import (
LeaveWorkspaceRequest,
LeaveWorkspaceUseCase,
RemoveMemberRequest,
RemoveMemberUseCase,
)
from packages.application.workspace.subscription_use_case import (
UpgradeSubscriptionUseCase,
CancelSubscriptionRequest,
CancelSubscriptionUseCase,
UpgradeSubscriptionRequest,
UpgradeSubscriptionResponse,
CancelSubscriptionUseCase,
CancelSubscriptionRequest,
UpgradeSubscriptionUseCase,
)
from packages.application.workspace.update_member_role_use_case import (
UpdateMemberRoleRequest,
UpdateMemberRoleResponse,
UpdateMemberRoleUseCase,
)
__all__ = [
@@ -1,14 +1,12 @@
"""
接受/拒绝邀请 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:
@@ -107,11 +105,14 @@ class AcceptInvitationUseCase:
invitation.accepted_at = datetime.now(timezone.utc)
self.workspace_invitation_repository.save(invitation)
return AcceptInvitationResponse(
return (
AcceptInvitationResponse(
workspace_id=workspace.id,
workspace_name=workspace.name,
role=existing_member.role,
), None
),
None,
)
# 9. 创建成员记录
member = WorkspaceMember(
@@ -131,11 +132,14 @@ class AcceptInvitationUseCase:
self.workspace_invitation_repository.save(invitation)
# 11. 返回响应
return AcceptInvitationResponse(
return (
AcceptInvitationResponse(
workspace_id=workspace.id,
workspace_name=workspace.name,
role=member.role,
), None
),
None,
)
except Exception as e:
return None, f"Failed to accept invitation: {str(e)}"
@@ -1,9 +1,10 @@
"""
创建 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
@@ -122,13 +123,16 @@ class CreateWorkspaceUseCase:
self.workspace_member_repository.save(owner_member)
# 8. 返回响应
return CreateWorkspaceResponse(
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
),
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,9 +9,9 @@ from uuid import uuid4
from packages.adapters.smtp import get_email_service
from packages.domain.entities import (
InvitationStatus,
WorkspaceInvitation,
WorkspaceMemberRole,
InvitationStatus,
)
@@ -114,7 +115,10 @@ 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. 检查被邀请人是否已经是成员
@@ -176,12 +180,15 @@ class InviteMemberUseCase:
print(f"Email service error: {e}")
# 11. 返回响应
return InviteMemberResponse(
return (
InviteMemberResponse(
invitation_id=invitation.id,
invitee_email=invitation.invitee_email,
role=invitation.role,
expires_at=invitation.expires_at,
), None
),
None,
)
except Exception as e:
return None, f"Failed to invite member: {str(e)}"
@@ -1,8 +1,9 @@
"""
获取成员列表 Use Case
"""
from typing import Optional, List
from datetime import datetime
from typing import List, Optional
class MemberInfo:
@@ -1,8 +1,9 @@
"""
获取工作空间列表和详情 Use Case
"""
from typing import Optional, List
from datetime import datetime
from typing import List, Optional
class WorkspaceInfo:
@@ -1,6 +1,7 @@
"""
移除成员 Use Case
"""
from typing import Optional
from packages.domain.entities import WorkspaceMemberRole
@@ -65,7 +66,10 @@ 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. 验证目标成员存在
@@ -85,8 +89,7 @@ class RemoveMemberUseCase:
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. 删除成员记录
@@ -152,7 +155,10 @@ class LeaveWorkspaceUseCase:
# 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)
@@ -1,8 +1,9 @@
"""
Subscription 管理 Use Case
"""
from typing import Optional
from datetime import datetime, timedelta, timezone
from typing import Optional
from packages.domain.entities import WorkspaceMemberRole
@@ -64,7 +65,9 @@ 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]]:
"""
执行升级订阅
@@ -110,7 +113,10 @@ class UpgradeSubscriptionUseCase:
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"
@@ -130,13 +136,16 @@ class UpgradeSubscriptionUseCase:
self.workspace_repository.save(workspace)
# 7. 返回响应
return UpgradeSubscriptionResponse(
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
),
None,
)
except Exception as e:
return None, f"Failed to upgrade subscription: {str(e)}"
@@ -1,6 +1,7 @@
"""
修改成员角色 Use Case
"""
from typing import Optional
from packages.domain.entities import WorkspaceMemberRole
@@ -74,7 +75,10 @@ class UpdateMemberRoleUseCase:
# 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)
@@ -89,7 +93,10 @@ 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. 验证目标成员存在
@@ -109,8 +116,7 @@ class UpdateMemberRoleUseCase:
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. 检查角色是否相同
@@ -123,11 +129,14 @@ class UpdateMemberRoleUseCase:
self.workspace_member_repository.save(target_member)
# 11. 返回响应
return UpdateMemberRoleResponse(
return (
UpdateMemberRoleResponse(
user_id=request.target_user_id,
old_role=old_role,
new_role=request.new_role,
), None
),
None,
)
except Exception as e:
return None, f"Failed to update member role: {str(e)}"
+6 -2
View File
@@ -1,6 +1,10 @@
"""Domain package for core business entities and rules."""
from .classification import AssetClassification, ClassificationJob, ClassificationJobStatus
from .classification import (
AssetClassification,
ClassificationJob,
ClassificationJobStatus,
)
from .entities import (
Asset,
AssetLibrary,
@@ -13,8 +17,8 @@ from .entities import (
User,
Workspace,
)
from .generation_task import GenerationTask, GenerationTaskStatus
from .generated_video import GeneratedVideo
from .generation_task import GenerationTask, GenerationTaskStatus
from .project_management import Milestone, Task, TaskIssue, TaskPriority, TaskStatus
__all__ = [
+7 -2
View File
@@ -5,7 +5,13 @@ services such as Redis session storage and SMTP email delivery live under
`packages.adapters` and should be injected into use cases.
"""
from packages.domain.auth.jwt_service import JWTConfig, JWTService, TokenType, jwt_service
from packages.domain.auth.email_service import EmailConfig, EmailService
from packages.domain.auth.jwt_service import (
JWTConfig,
JWTService,
TokenType,
jwt_service,
)
from packages.domain.auth.password_hasher import (
PasswordHasher,
PasswordValidator,
@@ -13,7 +19,6 @@ from packages.domain.auth.password_hasher import (
password_validator,
)
from packages.domain.auth.session_store import RedisConfig, SessionStore
from packages.domain.auth.email_service import EmailConfig, EmailService
__all__ = [
"JWTService",
+11 -26
View File
@@ -2,14 +2,17 @@
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"
@@ -19,6 +22,7 @@ class JWTConfig:
class TokenType:
"""Token 类型"""
ACCESS = "access"
REFRESH = "refresh"
@@ -34,7 +38,7 @@ class JWTService:
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
@@ -63,17 +67,9 @@ class JWTService:
if additional_claims:
payload.update(additional_claims)
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 create_refresh_token(
self,
user_id: str,
session_id: str
) -> str:
def create_refresh_token(self, user_id: str, session_id: str) -> str:
"""
创建 refresh_token
@@ -95,11 +91,7 @@ class JWTService:
"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]:
"""
@@ -116,11 +108,7 @@ class JWTService:
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")
@@ -182,10 +170,7 @@ class JWTService:
Token payload如果解码失败返回 None
"""
try:
return jwt.decode(
token,
options={"verify_signature": False}
)
return jwt.decode(token, options={"verify_signature": False})
except Exception:
return None
+9 -7
View File
@@ -2,9 +2,11 @@
密码哈希工具类
使用 bcrypt 安全存储密码
"""
import bcrypt
from typing import Optional
import bcrypt
class PasswordHasher:
"""密码哈希服务"""
@@ -39,14 +41,14 @@ class PasswordHasher:
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:
"""
@@ -63,8 +65,8 @@ class PasswordHasher:
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:
@@ -83,12 +85,12 @@ class PasswordHasher:
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
+1
View File
@@ -28,6 +28,7 @@ class ClassificationJobStatus(StrEnum):
class AssetClassification(StrEnum):
"""Asset classification categories."""
SCENIC = "scenic" # 风景
PRODUCT = "product" # 产品
PERSON = "person" # 人物
+4
View File
@@ -55,6 +55,7 @@ class Workspace:
class WorkspaceMemberRole(StrEnum):
"""工作空间成员角色"""
OWNER = "owner" # 所有者(创建者,唯一)
ADMIN = "admin" # 管理员(可管理成员和项目)
MEMBER = "member" # 成员(可创建和编辑项目)
@@ -63,6 +64,7 @@ class WorkspaceMemberRole(StrEnum):
class InvitationStatus(StrEnum):
"""邀请状态"""
PENDING = "pending" # 待处理
ACCEPTED = "accepted" # 已接受
DECLINED = "declined" # 已拒绝
@@ -72,6 +74,7 @@ class InvitationStatus(StrEnum):
@dataclass(slots=True)
class WorkspaceMember:
"""工作空间成员"""
id: str
workspace_id: str
user_id: str
@@ -83,6 +86,7 @@ class WorkspaceMember:
@dataclass(slots=True)
class WorkspaceInvitation:
"""工作空间邀请"""
id: str
workspace_id: str
inviter_user_id: str # 邀请人
+2
View File
@@ -2,7 +2,9 @@
权限验证辅助函数
用于检查用户在工作空间中的权限
"""
from typing import Optional
from packages.domain.entities import WorkspaceMemberRole
+6
View File
@@ -1,4 +1,5 @@
"""项目管理领域对象:任务、里程碑、项目阶段"""
from __future__ import annotations
from dataclasses import dataclass, field
@@ -9,6 +10,7 @@ from uuid import uuid4
class TaskStatus(StrEnum):
"""任务状态"""
PENDING = "pending" # 待开始
IN_PROGRESS = "in_progress" # 进行中
BLOCKED = "blocked" # 阻塞
@@ -18,6 +20,7 @@ class TaskStatus(StrEnum):
class TaskPriority(StrEnum):
"""任务优先级"""
LOW = "low"
MEDIUM = "medium"
HIGH = "high"
@@ -27,6 +30,7 @@ class TaskPriority(StrEnum):
@dataclass(slots=True)
class Task:
"""任务实体"""
id: str
project_id: str
workspace_id: str
@@ -127,6 +131,7 @@ class Task:
@dataclass(slots=True)
class Milestone:
"""里程碑实体"""
id: str
project_id: str
workspace_id: str
@@ -183,6 +188,7 @@ class Milestone:
@dataclass(slots=True)
class TaskIssue:
"""任务问题/卡点实体"""
id: str
task_id: str
project_id: str
+11 -8
View File
@@ -2,6 +2,7 @@
配额检查服务
用于检查工作空间是否超出配额限制
"""
from typing import Optional
@@ -38,7 +39,10 @@ class QuotaChecker:
# 检查是否超出配额(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
@@ -66,7 +70,10 @@ class QuotaChecker:
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
@@ -89,15 +96,11 @@ class QuotaChecker:
# 计算使用率
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 {
+5 -1
View File
@@ -3,7 +3,11 @@
from .asset_library_repository import AssetLibraryRepository
from .asset_repository import AssetRepository
from .ingest_job_repository import IngestJobRepository
from .project_management_repositories import MilestoneRepository, TaskIssueRepository, TaskRepository
from .project_management_repositories import (
MilestoneRepository,
TaskIssueRepository,
TaskRepository,
)
from .project_repository import ProjectRepository
__all__ = [
+4 -8
View File
@@ -6,14 +6,10 @@ from packages.domain import GeneratedVideo
class GeneratedVideoRepository(Protocol):
def create(self, video: GeneratedVideo) -> GeneratedVideo:
...
def create(self, video: GeneratedVideo) -> GeneratedVideo: ...
def get(self, video_id: str) -> GeneratedVideo | None:
...
def get(self, video_id: str) -> GeneratedVideo | None: ...
def list_by_project(self, project_id: str) -> list[GeneratedVideo]:
...
def list_by_project(self, project_id: str) -> list[GeneratedVideo]: ...
def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]:
...
def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]: ...
+4 -8
View File
@@ -6,14 +6,10 @@ from packages.domain import GenerationTask
class GenerationTaskRepository(Protocol):
def create(self, task: GenerationTask) -> GenerationTask:
...
def create(self, task: GenerationTask) -> GenerationTask: ...
def get(self, task_id: str) -> GenerationTask | None:
...
def get(self, task_id: str) -> GenerationTask | None: ...
def list_by_project(self, project_id: str) -> list[GenerationTask]:
...
def list_by_project(self, project_id: str) -> list[GenerationTask]: ...
def update(self, task: GenerationTask) -> GenerationTask:
...
def update(self, task: GenerationTask) -> GenerationTask: ...
@@ -1,4 +1,5 @@
"""项目管理 Repository 接口定义"""
from abc import ABC, abstractmethod
from packages.domain import Milestone, Task, TaskIssue
+2
View File
@@ -1,8 +1,10 @@
"""
Project 仓储接口
"""
from abc import ABC, abstractmethod
from typing import Optional
from packages.domain.entities import Project
+2
View File
@@ -1,8 +1,10 @@
"""
用户仓储接口
"""
from abc import ABC, abstractmethod
from typing import Optional
from packages.domain.entities import User
@@ -1,8 +1,10 @@
"""
WorkspaceInvitation 仓储接口
"""
from abc import ABC, abstractmethod
from typing import Optional
from packages.domain.entities import WorkspaceInvitation
@@ -1,8 +1,10 @@
"""
WorkspaceMember 仓储接口
"""
from abc import ABC, abstractmethod
from typing import Optional, List
from typing import List, Optional
from packages.domain.entities import WorkspaceMember

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