From 52ff2f80ad44e93cf062593f0f4e39d9c87c93ed Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Tue, 30 Jun 2026 17:23:08 +0800 Subject: [PATCH] style: apply black formatting to pass CI validation (#126) --- alembic/versions/007_add_editing_mode.py | 14 +- .../versions/008_add_video_dedup_fields.py | 16 +- .../versions/009_remove_workspace_concept.py | 33 +- alembic/versions/010_phase0_extensibility.py | 45 +-- alembic/versions/011_phase1_core_refactor.py | 9 +- .../versions/012_add_duplication_tables.py | 15 +- alembic/versions/013_add_recipe_tables.py | 9 +- alembic/versions/014_add_template_tables.py | 23 +- apps/api/app/api/routes/asset_libraries.py | 4 +- apps/api/app/api/routes/assets.py | 2 +- apps/api/app/api/routes/auth.py | 2 +- apps/api/app/api/routes/chunked_upload.py | 23 +- apps/api/app/api/routes/duplication.py | 27 +- apps/api/app/api/routes/generated_videos.py | 4 +- apps/api/app/api/routes/recipes.py | 6 +- apps/api/app/api/routes/subscription.py | 3 + apps/api/app/api/routes/templates.py | 16 +- apps/api/app/api/routes/titles.py | 1 + apps/api/app/api/routes/upload.py | 46 ++- apps/api/app/api/routes/voices.py | 1 + apps/api/app/config.py | 20 +- apps/api/app/core/database.py | 4 +- apps/api/app/core/storage.py | 3 - apps/api/app/dependencies.py | 3 +- apps/api/app/middleware/logging.py | 18 +- apps/api/app/middleware/prometheus_metrics.py | 2 +- apps/api/app/schemas/chunked_upload.py | 4 +- apps/api/app/schemas/dashboard.py | 2 + apps/api/app/schemas/duplication.py | 1 + apps/api/app/schemas/generation_task.py | 2 + apps/api/app/schemas/recipe.py | 4 +- apps/api/app/schemas/subscription.py | 15 +- apps/api/app/schemas/task_center.py | 2 + apps/api/app/schemas/template.py | 7 +- apps/api/app/schemas/title_library.py | 1 + apps/api/app/schemas/upload.py | 3 +- apps/api/app/schemas/voice_library.py | 1 + apps/worker/video_processing/dedup.py | 82 ++-- apps/worker/video_processing/editing_modes.py | 358 ++++++++++++++---- apps/worker/worker_app/tasks/__init__.py | 15 +- .../worker/worker_app/tasks/asset_analyzer.py | 104 ++--- .../worker/worker_app/tasks/classification.py | 3 +- apps/worker/worker_app/tasks/generation.py | 48 ++- apps/worker/worker_app/tasks/ingest.py | 6 +- .../worker_app/tasks/voice_extraction.py | 43 ++- fix_script.py | 6 +- fix_tracker_encoding.py | 180 ++++----- init_tracker_phase4.py | 170 +++++---- init_tracker_simple.py | 36 +- .../asset_library_repository.py | 20 +- .../sqlalchemy_impl/asset_repository.py | 33 +- .../sqlalchemy_impl/duplication_repository.py | 25 +- .../generation_task_repository.py | 6 +- packages/adapters/sqlalchemy_impl/models.py | 30 +- .../sqlalchemy_impl/project_repository.py | 19 +- .../sqlalchemy_impl/recipe_repository.py | 7 +- .../sqlalchemy_impl/template_repository.py | 5 +- .../title_library_repository.py | 49 ++- .../voice_library_repository.py | 49 ++- packages/adapters/sqlite_tracker/__init__.py | 1 - packages/application/auth/jwt_handler.py | 28 +- packages/application/auth/jwt_service.py | 24 +- packages/application/auth/login_use_case.py | 6 +- packages/application/auth/password_handler.py | 30 +- packages/application/duplication.py | 1 + packages/application/recipe/commands.py | 1 + packages/application/recipe/use_cases.py | 7 +- packages/application/template/commands.py | 1 + packages/application/template/use_cases.py | 36 +- .../application/title_library/__init__.py | 1 + .../application/title_library/commands.py | 1 + .../application/title_library/use_cases.py | 1 + .../application/voice_library/__init__.py | 1 + .../application/voice_library/commands.py | 1 + .../application/voice_library/use_cases.py | 1 + packages/domain/duplication.py | 1 + packages/domain/editing_mode.py | 9 +- packages/domain/entities.py | 7 - packages/domain/quota.py | 39 +- packages/domain/recipe.py | 3 + packages/domain/template.py | 4 + packages/domain/title_library.py | 2 + packages/domain/voice_library.py | 2 + packages/infrastructure/feature_flags.py | 7 +- packages/infrastructure/module_registry.py | 43 +-- packages/ports/duplication_repository.py | 1 + packages/ports/recipe_repository.py | 28 +- packages/ports/template_repository.py | 1 + packages/ports/title_library_repository.py | 19 +- packages/ports/voice_library_repository.py | 19 +- scripts/smoke_test.py | 85 +++-- tests/integration/test_auth.py | 1 - .../test_duplication_upload_error_handling.py | 95 +++-- tests/integration/test_generation_pipeline.py | 4 +- tests/integration/test_projects.py | 1 - tests/integration/test_subscription_api.py | 160 +++++--- tests/unit/test_asset_diagnosis.py | 4 +- tests/unit/test_auth_middleware.py | 2 - tests/unit/test_chunked_upload_routes.py | 2 +- tests/unit/test_config_oss.py | 1 + tests/unit/test_form_upload_routes.py | 2 +- tests/unit/test_recipe_use_cases.py | 1 + tests/unit/test_template_use_cases.py | 75 +++- tests/unit/test_title_library_use_cases.py | 12 +- tests/unit/test_upload_routes.py | 2 +- tests/unit/test_voice_library_use_cases.py | 22 +- update_tracker.py | 47 ++- 107 files changed, 1517 insertions(+), 1010 deletions(-) diff --git a/alembic/versions/007_add_editing_mode.py b/alembic/versions/007_add_editing_mode.py index 6655a6783..3ab6884cf 100755 --- a/alembic/versions/007_add_editing_mode.py +++ b/alembic/versions/007_add_editing_mode.py @@ -4,25 +4,25 @@ Revision ID: 007 Revises: 006 Create Date: 2026-06-26 """ + from alembic import op import sqlalchemy as sa # revision identifiers -revision = '007' -down_revision = '006' +revision = "007" +down_revision = "006" branch_labels = None depends_on = None def upgrade() -> None: op.add_column( - 'generation_tasks', - sa.Column('editing_mode', sa.String(20), nullable=False, server_default='one_take') + "generation_tasks", sa.Column("editing_mode", sa.String(20), nullable=False, server_default="one_take") ) # 添加索引以支持查询 - op.create_index('ix_generation_tasks_editing_mode', 'generation_tasks', ['editing_mode']) + op.create_index("ix_generation_tasks_editing_mode", "generation_tasks", ["editing_mode"]) def downgrade() -> None: - op.drop_index('ix_generation_tasks_editing_mode', table_name='generation_tasks') - op.drop_column('generation_tasks', 'editing_mode') + op.drop_index("ix_generation_tasks_editing_mode", table_name="generation_tasks") + op.drop_column("generation_tasks", "editing_mode") diff --git a/alembic/versions/008_add_video_dedup_fields.py b/alembic/versions/008_add_video_dedup_fields.py index e304cff44..45b468615 100644 --- a/alembic/versions/008_add_video_dedup_fields.py +++ b/alembic/versions/008_add_video_dedup_fields.py @@ -4,6 +4,7 @@ Revision ID: 008 Revises: 007 Create Date: 2024-06-26 """ + import sqlalchemy as sa from alembic import op @@ -16,20 +17,11 @@ depends_on = None def upgrade() -> None: # Add video_fingerprint column as JSON text - op.add_column( - "generated_videos", - sa.Column("video_fingerprint", sa.Text(), nullable=True) - ) + op.add_column("generated_videos", sa.Column("video_fingerprint", sa.Text(), nullable=True)) # Add is_duplicate column - op.add_column( - "generated_videos", - sa.Column("is_duplicate", sa.Boolean(), nullable=False, server_default="false") - ) + op.add_column("generated_videos", sa.Column("is_duplicate", sa.Boolean(), nullable=False, server_default="false")) # Add duplicate_of column for tracking original video - op.add_column( - "generated_videos", - sa.Column("duplicate_of", sa.String(32), nullable=True) - ) + op.add_column("generated_videos", sa.Column("duplicate_of", sa.String(32), nullable=True)) def downgrade() -> None: diff --git a/alembic/versions/009_remove_workspace_concept.py b/alembic/versions/009_remove_workspace_concept.py index c490752c5..f41a4bd26 100644 --- a/alembic/versions/009_remove_workspace_concept.py +++ b/alembic/versions/009_remove_workspace_concept.py @@ -11,6 +11,7 @@ This migration: 4. Removes workspace_id from all tables that had it 5. Drops workspace-related tables: workspaces, workspace_members, workspace_invitations """ + from alembic import op import sqlalchemy as sa from sqlalchemy import text @@ -24,7 +25,7 @@ depends_on = None def upgrade() -> None: conn = op.get_bind() - + # Step 1: Add subscription/quota fields to users table conn.execute(text(""" ALTER TABLE users @@ -50,7 +51,7 @@ def upgrade() -> None: ALTER TABLE users ADD COLUMN IF NOT EXISTS used_storage_gb FLOAT NOT NULL DEFAULT 0 """)) - + # Step 2: Copy subscription data from workspaces to users conn.execute(text(""" UPDATE users SET @@ -63,7 +64,7 @@ def upgrade() -> None: FROM workspaces w WHERE w.owner_user_id = users.id """)) - + # Step 3: Add owner_user_id and shared_users to projects table conn.execute(text(""" ALTER TABLE projects @@ -73,7 +74,7 @@ def upgrade() -> None: ALTER TABLE projects ADD COLUMN IF NOT EXISTS shared_users JSON """)) - + # Step 4: Migrate workspace_id to owner_user_id (from workspace_members where role=owner) conn.execute(text(""" UPDATE projects SET @@ -82,13 +83,13 @@ def upgrade() -> None: WHERE wm.workspace_id = projects.workspace_id AND wm.role = 'owner' """)) - + # Set shared_users to empty array for all projects conn.execute(text(""" UPDATE projects SET shared_users = '[]'::json WHERE shared_users IS NULL """)) - + # Step 5: Remove workspace_id from all tables tables_with_workspace_id = [ "asset_libraries", @@ -104,12 +105,12 @@ def upgrade() -> None: "tasks", "task_issues", ] - + for table in tables_with_workspace_id: conn.execute(text(f""" ALTER TABLE {table} DROP COLUMN IF EXISTS workspace_id """)) - + # Step 6: Drop workspace-related tables conn.execute(text(""" DROP TABLE IF EXISTS workspace_invitations @@ -120,7 +121,7 @@ def upgrade() -> None: conn.execute(text(""" DROP TABLE IF EXISTS workspaces """)) - + # Step 7: Drop workspace_id from projects table conn.execute(text(""" ALTER TABLE projects DROP COLUMN IF EXISTS workspace_id @@ -129,7 +130,7 @@ def upgrade() -> None: def downgrade() -> None: conn = op.get_bind() - + # Add back workspace tables (simplified - in real scenario would need full recreation) conn.execute(text(""" CREATE TABLE IF NOT EXISTS workspaces ( @@ -145,7 +146,7 @@ def downgrade() -> None: created_at TIMESTAMP NOT NULL DEFAULT NOW() ) """)) - + conn.execute(text(""" CREATE TABLE IF NOT EXISTS workspace_members ( id VARCHAR(36) PRIMARY KEY, @@ -157,7 +158,7 @@ def downgrade() -> None: UNIQUE(workspace_id, user_id) ) """)) - + conn.execute(text(""" CREATE TABLE IF NOT EXISTS workspace_invitations ( id VARCHAR(36) PRIMARY KEY, @@ -172,12 +173,12 @@ def downgrade() -> None: created_at TIMESTAMP NOT NULL DEFAULT NOW() ) """)) - + # Add back workspace_id column to projects conn.execute(text(""" ALTER TABLE projects ADD COLUMN workspace_id VARCHAR(32) """)) - + # Add back workspace_id columns to other tables tables_with_workspace_id = [ "asset_libraries", @@ -193,11 +194,11 @@ def downgrade() -> None: "tasks", "task_issues", ] - + for table in tables_with_workspace_id: conn.execute(text(f""" ALTER TABLE {table} ADD COLUMN workspace_id VARCHAR(36) """)) - + # Note: This downgrade is incomplete - projects.owner_user_id data would need to be # converted back to workspace_ids, which requires reconstructing workspace records. diff --git a/alembic/versions/010_phase0_extensibility.py b/alembic/versions/010_phase0_extensibility.py index 512e058d4..66096b032 100644 --- a/alembic/versions/010_phase0_extensibility.py +++ b/alembic/versions/010_phase0_extensibility.py @@ -10,6 +10,7 @@ This migration: 2. Creates title_libraries table (独立标题库,支持跨项目复用) 3. Creates voice_libraries table (配音库,支持 AI 配音管理) """ + from alembic import op import sqlalchemy as sa @@ -25,21 +26,11 @@ def upgrade() -> None: # ── 1. Add metadata JSONB to existing tables ── - conn.execute(sa.text( - "ALTER TABLE projects ADD COLUMN IF NOT EXISTS metadata JSONB NOT NULL DEFAULT '{}'" - )) - conn.execute(sa.text( - "ALTER TABLE asset_libraries ADD COLUMN IF NOT EXISTS metadata JSONB NOT NULL DEFAULT '{}'" - )) - conn.execute(sa.text( - "ALTER TABLE assets ADD COLUMN IF NOT EXISTS metadata JSONB NOT NULL DEFAULT '{}'" - )) - conn.execute(sa.text( - "ALTER TABLE edit_templates ADD COLUMN IF NOT EXISTS metadata JSONB NOT NULL DEFAULT '{}'" - )) - conn.execute(sa.text( - "ALTER TABLE generation_tasks ADD COLUMN IF NOT EXISTS metadata JSONB NOT NULL DEFAULT '{}'" - )) + conn.execute(sa.text("ALTER TABLE projects ADD COLUMN IF NOT EXISTS metadata JSONB NOT NULL DEFAULT '{}'")) + conn.execute(sa.text("ALTER TABLE asset_libraries ADD COLUMN IF NOT EXISTS metadata JSONB NOT NULL DEFAULT '{}'")) + conn.execute(sa.text("ALTER TABLE assets ADD COLUMN IF NOT EXISTS metadata JSONB NOT NULL DEFAULT '{}'")) + conn.execute(sa.text("ALTER TABLE edit_templates ADD COLUMN IF NOT EXISTS metadata JSONB NOT NULL DEFAULT '{}'")) + conn.execute(sa.text("ALTER TABLE generation_tasks ADD COLUMN IF NOT EXISTS metadata JSONB NOT NULL DEFAULT '{}'")) # ── 2. Create title_libraries table ── @@ -59,15 +50,9 @@ def upgrade() -> None: updated_at TIMESTAMP NOT NULL DEFAULT NOW() ) """)) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_title_libraries_user_id ON title_libraries(user_id)" - )) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_title_libraries_category ON title_libraries(category)" - )) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_title_libraries_is_active ON title_libraries(is_active)" - )) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_title_libraries_user_id ON title_libraries(user_id)")) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_title_libraries_category ON title_libraries(category)")) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_title_libraries_is_active ON title_libraries(is_active)")) # ── 3. Create voice_libraries table ── @@ -91,15 +76,9 @@ def upgrade() -> None: updated_at TIMESTAMP NOT NULL DEFAULT NOW() ) """)) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_voice_libraries_user_id ON voice_libraries(user_id)" - )) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_voice_libraries_project_id ON voice_libraries(project_id)" - )) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_voice_libraries_status ON voice_libraries(status)" - )) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_voice_libraries_user_id ON voice_libraries(user_id)")) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_voice_libraries_project_id ON voice_libraries(project_id)")) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_voice_libraries_status ON voice_libraries(status)")) def downgrade() -> None: diff --git a/alembic/versions/011_phase1_core_refactor.py b/alembic/versions/011_phase1_core_refactor.py index 52e64c875..5d337101e 100644 --- a/alembic/versions/011_phase1_core_refactor.py +++ b/alembic/versions/011_phase1_core_refactor.py @@ -14,6 +14,7 @@ This migration: - edit_plan_clips (编辑计划片段) 2. Removes edit_plan_id column from generation_tasks table """ + from alembic import op import sqlalchemy as sa @@ -39,9 +40,7 @@ def upgrade() -> None: # ── 2. Remove edit_plan_id from generation_tasks ── - conn.execute(sa.text( - "ALTER TABLE generation_tasks DROP COLUMN IF EXISTS edit_plan_id" - )) + conn.execute(sa.text("ALTER TABLE generation_tasks DROP COLUMN IF EXISTS edit_plan_id")) def downgrade() -> None: @@ -49,9 +48,7 @@ def downgrade() -> None: # ── 1. Re-add edit_plan_id to generation_tasks ── - conn.execute(sa.text( - "ALTER TABLE generation_tasks ADD COLUMN IF NOT EXISTS edit_plan_id VARCHAR(32)" - )) + conn.execute(sa.text("ALTER TABLE generation_tasks ADD COLUMN IF NOT EXISTS edit_plan_id VARCHAR(32)")) # ── 2. Recreate deprecated tables (basic structure) ── diff --git a/alembic/versions/012_add_duplication_tables.py b/alembic/versions/012_add_duplication_tables.py index 8940b0e7c..97a4cce43 100644 --- a/alembic/versions/012_add_duplication_tables.py +++ b/alembic/versions/012_add_duplication_tables.py @@ -8,6 +8,7 @@ This migration creates two new tables: 1. duplication_records — 查重记录主表 2. duplication_segments — 重复片段详情表 """ + from alembic import op import sqlalchemy as sa @@ -40,12 +41,8 @@ def upgrade() -> None: updated_at TIMESTAMP NOT NULL DEFAULT NOW() ) """)) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_duplication_records_user_id ON duplication_records(user_id)" - )) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_duplication_records_status ON duplication_records(status)" - )) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_duplication_records_user_id ON duplication_records(user_id)")) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_duplication_records_status ON duplication_records(status)")) # ── 2. Create duplication_segments table ── @@ -62,9 +59,9 @@ def upgrade() -> None: similarity FLOAT NOT NULL ) """)) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_duplication_segments_record_id ON duplication_segments(record_id)" - )) + conn.execute( + sa.text("CREATE INDEX IF NOT EXISTS ix_duplication_segments_record_id ON duplication_segments(record_id)") + ) def downgrade() -> None: diff --git a/alembic/versions/013_add_recipe_tables.py b/alembic/versions/013_add_recipe_tables.py index bd1357d27..d74f077c0 100644 --- a/alembic/versions/013_add_recipe_tables.py +++ b/alembic/versions/013_add_recipe_tables.py @@ -8,6 +8,7 @@ This migration creates two new tables: 1. recipes — 配方主表 2. recipe_items — 配方素材项表 """ + from alembic import op import sqlalchemy as sa @@ -37,9 +38,7 @@ def upgrade() -> None: updated_at TIMESTAMP NOT NULL DEFAULT NOW() ) """)) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_recipes_user_id ON recipes(user_id)" - )) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_recipes_user_id ON recipes(user_id)")) # ── 2. Create recipe_items table ── @@ -53,9 +52,7 @@ def upgrade() -> None: metadata JSONB NOT NULL DEFAULT '{}' ) """)) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_recipe_items_recipe_id ON recipe_items(recipe_id)" - )) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_recipe_items_recipe_id ON recipe_items(recipe_id)")) def downgrade() -> None: diff --git a/alembic/versions/014_add_template_tables.py b/alembic/versions/014_add_template_tables.py index 3ae2ee4f8..619835f99 100644 --- a/alembic/versions/014_add_template_tables.py +++ b/alembic/versions/014_add_template_tables.py @@ -9,6 +9,7 @@ This migration creates three new tables: 2. template_segments — 模板片段表 3. template_categories — 模板分类表 """ + from alembic import op import sqlalchemy as sa @@ -40,12 +41,8 @@ def upgrade() -> None: updated_at TIMESTAMP NOT NULL DEFAULT NOW() ) """)) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_templates_user_id ON templates(user_id)" - )) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_templates_mode ON templates(mode)" - )) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_templates_user_id ON templates(user_id)")) + conn.execute(sa.text("CREATE INDEX IF NOT EXISTS ix_templates_mode ON templates(mode)")) # ── 2. Create template_segments table ── conn.execute(sa.text(""" @@ -60,10 +57,9 @@ def upgrade() -> None: updated_at TIMESTAMP NOT NULL DEFAULT NOW() ) """)) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_template_segments_template_id " - "ON template_segments(template_id)" - )) + conn.execute( + sa.text("CREATE INDEX IF NOT EXISTS ix_template_segments_template_id " "ON template_segments(template_id)") + ) # ── 3. Create template_categories table ── conn.execute(sa.text(""" @@ -74,10 +70,9 @@ def upgrade() -> None: created_at TIMESTAMP NOT NULL DEFAULT NOW() ) """)) - conn.execute(sa.text( - "CREATE INDEX IF NOT EXISTS ix_template_categories_user_id " - "ON template_categories(user_id)" - )) + conn.execute( + sa.text("CREATE INDEX IF NOT EXISTS ix_template_categories_user_id " "ON template_categories(user_id)") + ) def downgrade() -> None: diff --git a/apps/api/app/api/routes/asset_libraries.py b/apps/api/app/api/routes/asset_libraries.py index c8d25228a..af09c629f 100644 --- a/apps/api/app/api/routes/asset_libraries.py +++ b/apps/api/app/api/routes/asset_libraries.py @@ -49,7 +49,7 @@ def list_asset_libraries( ) -> ListAssetLibrariesResponse: user_id = authenticated_user.user.id use_case = ListAssetLibrariesUseCase(asset_library_repository) - + if project_id: # If project_id provided, check access and filter by project project = GetProjectUseCase(project_repository).execute(project_id) @@ -65,7 +65,7 @@ def list_asset_libraries( for proj in accessible_projects: all_items.extend(use_case.execute(proj.id)) items = all_items - + return ListAssetLibrariesResponse(items=[_to_asset_library_response(item) for item in items]) diff --git a/apps/api/app/api/routes/assets.py b/apps/api/app/api/routes/assets.py index aa37bfe99..2a8181916 100644 --- a/apps/api/app/api/routes/assets.py +++ b/apps/api/app/api/routes/assets.py @@ -106,7 +106,7 @@ def create_asset( raise HTTPException(status_code=404, detail=f"Project {request.project_id} not found") if not project.can_access(authenticated_user.user.id): raise HTTPException(status_code=403, detail="Access denied to project") - + library = asset_library_repository.get(request.library_id) if library is None or library.project_id != request.project_id: raise HTTPException(status_code=404, detail=f"AssetLibrary {request.library_id} not found") diff --git a/apps/api/app/api/routes/auth.py b/apps/api/app/api/routes/auth.py index 1ded9dd1a..92d7394ab 100755 --- a/apps/api/app/api/routes/auth.py +++ b/apps/api/app/api/routes/auth.py @@ -55,6 +55,7 @@ class LoginRequest(BaseModel): class RefreshRequest(BaseModel): refresh_token: str + class LoginResponse(BaseModel): access_token: str refresh_token: str @@ -173,7 +174,6 @@ async def refresh( ) - def _verify_email_token(token: str, user_repository: UserRepository) -> MessageResponse: success, error = VerifyEmailUseCase(user_repository=user_repository).execute(VerifyEmailRequest(token=token)) if not success: diff --git a/apps/api/app/api/routes/chunked_upload.py b/apps/api/app/api/routes/chunked_upload.py index 4e06b83d7..ddb094654 100644 --- a/apps/api/app/api/routes/chunked_upload.py +++ b/apps/api/app/api/routes/chunked_upload.py @@ -44,9 +44,18 @@ CHUNK_EXPIRY_HOURS = 24 # Allowed file types (consistent with existing upload.py) ALLOWED_MIME_TYPES = { - "image/jpeg", "image/png", "image/gif", "image/webp", - "video/mp4", "video/quicktime", "video/x-msvideo", "video/webm", - "audio/mpeg", "audio/wav", "audio/ogg", "audio/mp3", + "image/jpeg", + "image/png", + "image/gif", + "image/webp", + "video/mp4", + "video/quicktime", + "video/x-msvideo", + "video/webm", + "audio/mpeg", + "audio/wav", + "audio/ogg", + "audio/mp3", } # Chunk storage root directory @@ -67,13 +76,13 @@ def _atomic_check_and_record(upload_id: str, chunk_index: int) -> bool: """ Atomically check if chunk is uploaded and record if not. Uses file locking to prevent race conditions. - + Returns: True if chunk was newly recorded, False if already exists """ meta_path = _get_upload_meta_path(upload_id) CHUNK_STORAGE_ROOT.mkdir(parents=True, exist_ok=True) - + with open(meta_path, "r+", encoding="utf-8") as f: fcntl.flock(f.fileno(), fcntl.LOCK_EX) try: @@ -128,15 +137,17 @@ def _validate_file_type(content: bytes, filename: str) -> str: """Validate file type""" try: import magic + detected_mime = magic.from_buffer(content, mime=True) except ImportError: import mimetypes + detected_mime = mimetypes.guess_type(filename)[0] or "application/octet-stream" if detected_mime not in ALLOWED_MIME_TYPES: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, - detail=f"Unsupported file type: {detected_mime}. Allowed types: {', '.join(sorted(ALLOWED_MIME_TYPES))}" + detail=f"Unsupported file type: {detected_mime}. Allowed types: {', '.join(sorted(ALLOWED_MIME_TYPES))}", ) return detected_mime diff --git a/apps/api/app/api/routes/duplication.py b/apps/api/app/api/routes/duplication.py index 5507f65d9..4e4616e7e 100644 --- a/apps/api/app/api/routes/duplication.py +++ b/apps/api/app/api/routes/duplication.py @@ -1,4 +1,5 @@ """查重 API 路由。""" + from __future__ import annotations import logging @@ -31,10 +32,17 @@ logger = logging.getLogger(__name__) router = APIRouter() # 查重功能只接受视频文件 -ALLOWED_VIDEO_MIME_TYPES = frozenset({ - "video/mp4", "video/mpeg", "video/quicktime", "video/x-msvideo", - "video/webm", "video/x-matroska", "video/3gpp", -}) +ALLOWED_VIDEO_MIME_TYPES = frozenset( + { + "video/mp4", + "video/mpeg", + "video/quicktime", + "video/x-msvideo", + "video/webm", + "video/x-matroska", + "video/3gpp", + } +) def _validate_video_mime_type(content_type: str | None) -> str: @@ -44,16 +52,16 @@ def _validate_video_mime_type(content_type: str | None) -> str: status_code=status.HTTP_400_BAD_REQUEST, detail="Content-Type header is required", ) - + # 处理带参数的类型,如 "video/mp4; charset=utf-8" base_type = content_type.split(";")[0].strip().lower() - + if base_type not in ALLOWED_VIDEO_MIME_TYPES: raise HTTPException( status_code=status.HTTP_415_UNSUPPORTED_MEDIA_TYPE, detail=f"只支持视频文件。支持的类型: mp4, mpeg, mov, avi, webm, mkv, 3gp", ) - + return base_type @@ -117,9 +125,10 @@ async def upload_for_duplication( # P0-2: 验证文件大小(参考 OSS_DIRECT_UPLOAD_MAX_MB) from app.config import get_settings + settings = get_settings() max_size_bytes = settings.OSS_DIRECT_UPLOAD_MAX_MB * 1024 * 1024 - + # 先检查 Content-Length header(如果可用) if file.size is not None and file.size > max_size_bytes: raise HTTPException( @@ -135,7 +144,7 @@ async def upload_for_duplication( try: content = await file.read() file_size = len(content) - + # 再次检查实际文件大小 if file_size > max_size_bytes: raise HTTPException( diff --git a/apps/api/app/api/routes/generated_videos.py b/apps/api/app/api/routes/generated_videos.py index 30a1b6aa8..d9f72c031 100644 --- a/apps/api/app/api/routes/generated_videos.py +++ b/apps/api/app/api/routes/generated_videos.py @@ -50,7 +50,7 @@ def list_generated_videos( ) -> ListGeneratedVideosResponse: user_id = authenticated_user.user.id use_case = ListGeneratedVideosUseCase(generated_video_repository) - + if project_id: # If project_id provided, check access and filter by project project = project_repository.find_by_id(project_id) @@ -64,7 +64,7 @@ def list_generated_videos( for proj in accessible_projects: all_items.extend(use_case.execute(proj.id)) items = all_items - + # Generate download URLs for each video responses = [] for item in items: diff --git a/apps/api/app/api/routes/recipes.py b/apps/api/app/api/routes/recipes.py index f5c7d97e9..f00daf11d 100644 --- a/apps/api/app/api/routes/recipes.py +++ b/apps/api/app/api/routes/recipes.py @@ -1,4 +1,5 @@ """Recipe CRUD + use routes.""" + from __future__ import annotations from typing import List @@ -209,8 +210,5 @@ def use_recipe( return UseRecipeResponse( recipe=_to_response(result.recipe), - warnings=[ - {"item_type": w.item_type, "item_id": w.item_id, "position": w.position} - for w in result.warnings - ], + warnings=[{"item_type": w.item_type, "item_id": w.item_id, "position": w.position} for w in result.warnings], ) diff --git a/apps/api/app/api/routes/subscription.py b/apps/api/app/api/routes/subscription.py index 5cd0c7525..48b9d310b 100644 --- a/apps/api/app/api/routes/subscription.py +++ b/apps/api/app/api/routes/subscription.py @@ -1,4 +1,5 @@ """Subscription management API routes.""" + from __future__ import annotations from dataclasses import replace @@ -34,6 +35,7 @@ PLAN_QUOTAS = { # ============ Helper Functions ============ + def _get_plan_name(plan_id: str) -> str: """获取套餐显示名称""" plan_names = { @@ -86,6 +88,7 @@ def _build_subscription_info(user: AuthenticatedUser) -> SubscriptionInfo: # ============ API Endpoints ============ + @router.get("/current", response_model=SubscriptionInfo) async def get_current_subscription( current_user: AuthenticatedUser = Depends(get_current_user), diff --git a/apps/api/app/api/routes/templates.py b/apps/api/app/api/routes/templates.py index 2d6b6afae..7b8849bdc 100644 --- a/apps/api/app/api/routes/templates.py +++ b/apps/api/app/api/routes/templates.py @@ -1,4 +1,5 @@ """Template CRUD + generate + category routes.""" + from __future__ import annotations from typing import List @@ -234,10 +235,7 @@ def validate_template( return ValidateTemplateResponse( template=_to_response(result.template), - warnings=[ - GenerateWarningResponse(code=w.code, message=w.message, details=w.details) - for w in result.warnings - ], + warnings=[GenerateWarningResponse(code=w.code, message=w.message, details=w.details) for w in result.warnings], ) @@ -253,10 +251,7 @@ def list_categories( use_case = ListCategoriesUseCase(template_repository) categories = use_case.execute(user_id) return ListCategoriesResponse( - items=[ - CategoryResponse(id=c.id, user_id=c.user_id, name=c.name, created_at=c.created_at) - for c in categories - ], + items=[CategoryResponse(id=c.id, user_id=c.user_id, name=c.name, created_at=c.created_at) for c in categories], ) @@ -271,7 +266,10 @@ def create_category( use_case = CreateCategoryUseCase(template_repository) category = use_case.execute(command) return CategoryResponse( - id=category.id, user_id=category.user_id, name=category.name, created_at=category.created_at, + id=category.id, + user_id=category.user_id, + name=category.name, + created_at=category.created_at, ) diff --git a/apps/api/app/api/routes/titles.py b/apps/api/app/api/routes/titles.py index 7ed4b6285..634668e8e 100644 --- a/apps/api/app/api/routes/titles.py +++ b/apps/api/app/api/routes/titles.py @@ -1,4 +1,5 @@ """Title library CRUD routes.""" + from __future__ import annotations from typing import Optional diff --git a/apps/api/app/api/routes/upload.py b/apps/api/app/api/routes/upload.py index c11f918f7..65228168b 100644 --- a/apps/api/app/api/routes/upload.py +++ b/apps/api/app/api/routes/upload.py @@ -30,17 +30,35 @@ logger = logging.getLogger(__name__) router = APIRouter() # 允许上传的文件 MIME 类型 -ALLOWED_MIME_TYPES = frozenset({ - # 视频 - "video/mp4", "video/mpeg", "video/quicktime", "video/x-msvideo", - "video/webm", "video/x-matroska", "video/3gpp", - # 音频 - "audio/mpeg", "audio/wav", "audio/ogg", "audio/flac", "audio/aac", - "audio/mp3", "audio/x-m4a", "audio/webm", - # 图片 - "image/jpeg", "image/png", "image/gif", "image/webp", "image/bmp", - "image/svg+xml", "image/tiff", -}) +ALLOWED_MIME_TYPES = frozenset( + { + # 视频 + "video/mp4", + "video/mpeg", + "video/quicktime", + "video/x-msvideo", + "video/webm", + "video/x-matroska", + "video/3gpp", + # 音频 + "audio/mpeg", + "audio/wav", + "audio/ogg", + "audio/flac", + "audio/aac", + "audio/mp3", + "audio/x-m4a", + "audio/webm", + # 图片 + "image/jpeg", + "image/png", + "image/gif", + "image/webp", + "image/bmp", + "image/svg+xml", + "image/tiff", + } +) def _validate_mime_type(content_type: str | None) -> str: @@ -50,16 +68,16 @@ def _validate_mime_type(content_type: str | None) -> str: status_code=status.HTTP_400_BAD_REQUEST, detail="Content-Type header is required", ) - + # 处理带参数的类型,如 "video/mp4; charset=utf-8" base_type = content_type.split(";")[0].strip().lower() - + if base_type not in ALLOWED_MIME_TYPES: raise HTTPException( status_code=status.HTTP_415_UNSUPPORTED_MEDIA_TYPE, detail=f"File type '{base_type}' is not supported. Allowed types: video, audio, and image files.", ) - + return base_type diff --git a/apps/api/app/api/routes/voices.py b/apps/api/app/api/routes/voices.py index ed210e81b..1512c6867 100644 --- a/apps/api/app/api/routes/voices.py +++ b/apps/api/app/api/routes/voices.py @@ -1,4 +1,5 @@ """Voice library CRUD routes.""" + from __future__ import annotations from typing import Optional diff --git a/apps/api/app/config.py b/apps/api/app/config.py index 778492c04..da2b61152 100755 --- a/apps/api/app/config.py +++ b/apps/api/app/config.py @@ -21,9 +21,7 @@ class Settings(BaseSettings): API_PORT: int = 8000 API_PREFIX: str = "/api/v1" - DATABASE_URL: str = ( - "postgresql+psycopg://postgres:postgres@localhost:5432/xiaoxia_saas" - ) + DATABASE_URL: str = "postgresql+psycopg://postgres:postgres@localhost:5432/xiaoxia_saas" DATABASE_POOL_SIZE: int = 20 DATABASE_MAX_OVERFLOW: int = 10 # 调整为合理值:pool_size(20) + max_overflow(10) = 最大30连接 DATABASE_POOL_TIMEOUT: int = 30 @@ -48,8 +46,7 @@ class Settings(BaseSettings): def validate_jwt_secret_key(cls, v): if v is None or v == "": raise ValueError( - "JWT_SECRET_KEY must be set via environment variable. " - "Do not use default value in production!" + "JWT_SECRET_KEY must be set via environment variable. " "Do not use default value in production!" ) # Block known insecure default values insecure_defaults = [ @@ -61,8 +58,7 @@ class Settings(BaseSettings): ] if v.lower() in [d.lower() for d in insecure_defaults]: raise ValueError( - f"JWT_SECRET_KEY '{v}' is insecure. " - "Please set a strong random secret via environment variable." + f"JWT_SECRET_KEY '{v}' is insecure. " "Please set a strong random secret via environment variable." ) return v @@ -90,9 +86,7 @@ class Settings(BaseSettings): OSS_DIRECT_UPLOAD_EXPIRE_SECONDS: int = 900 LOG_LEVEL: str = "INFO" - CORS_ORIGINS_RAW: str = ( - "http://localhost:3000,http://localhost:5173,http://localhost:8000" - ) + CORS_ORIGINS_RAW: str = "http://localhost:3000,http://localhost:5173,http://localhost:8000" model_config = SettingsConfigDict( env_file=".env", @@ -103,11 +97,7 @@ class Settings(BaseSettings): @property def CORS_ORIGINS(self) -> list[str]: - return [ - origin.strip() - for origin in self.CORS_ORIGINS_RAW.split(",") - if origin.strip() - ] + return [origin.strip() for origin in self.CORS_ORIGINS_RAW.split(",") if origin.strip()] @property def database_url(self) -> str: diff --git a/apps/api/app/core/database.py b/apps/api/app/core/database.py index 94a26349d..82149f15c 100644 --- a/apps/api/app/core/database.py +++ b/apps/api/app/core/database.py @@ -2,6 +2,7 @@ 统一使用 app.config 中的数据库配置,移除重复的 DatabaseSettings。 """ + from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker, Session from contextlib import contextmanager @@ -9,7 +10,6 @@ from typing import Generator from app.config import settings - engine = create_engine( settings.database_url, pool_size=settings.DATABASE_POOL_SIZE, @@ -33,7 +33,7 @@ def get_db() -> Generator[Session, None, None]: @contextmanager def get_db_context() -> Generator[Session, None, None]: """Context manager for database sessions. - + Usage: with get_db_context() as db: db.query(Model).all() diff --git a/apps/api/app/core/storage.py b/apps/api/app/core/storage.py index 094c3bea8..cf6997b61 100644 --- a/apps/api/app/core/storage.py +++ b/apps/api/app/core/storage.py @@ -232,8 +232,6 @@ class OSSStorageService: return self.bucket.object_exists(storage_key) - - _storage_service = None @@ -243,4 +241,3 @@ def get_storage_service() -> OSSStorageService: if _storage_service is None: _storage_service = OSSStorageService() return _storage_service - diff --git a/apps/api/app/dependencies.py b/apps/api/app/dependencies.py index 79bcdffdc..5c3b4a7bb 100644 --- a/apps/api/app/dependencies.py +++ b/apps/api/app/dependencies.py @@ -3,6 +3,7 @@ All repository and service factories are defined here as FastAPI dependencies, ensuring proper lifecycle management and testability. """ + from __future__ import annotations import redis @@ -124,7 +125,6 @@ def get_project_repository( return SQLAlchemyProjectRepository(session) - def get_user_repository( session: Session = Depends(get_db_session), ) -> UserRepository: @@ -156,6 +156,7 @@ def get_auth_email_service() -> NoopEmailService | EmailService: enabled=True, ) + def get_title_library_repository( session: Session = Depends(get_db_session), ) -> SQLAlchemyTitleLibraryRepository: diff --git a/apps/api/app/middleware/logging.py b/apps/api/app/middleware/logging.py index c0095a34d..f17bb1ea7 100644 --- a/apps/api/app/middleware/logging.py +++ b/apps/api/app/middleware/logging.py @@ -69,9 +69,7 @@ class RequestLoggingMiddleware(BaseHTTPMiddleware): # 记录请求信息(不包含敏感参数) if safe_query: - logger.info( - f"Request: {request.method} {request.url.path}?{safe_query}" # noqa: E501 - ) + logger.info(f"Request: {request.method} {request.url.path}?{safe_query}") # noqa: E501 else: logger.info(f"Request: {request.method} {request.url.path}") @@ -83,8 +81,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" ) # 添加处理时间到响应头 @@ -111,9 +108,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 ] # 计算请求次数 @@ -128,8 +123,7 @@ class RateLimitMiddleware(BaseHTTPMiddleware): "error": { "code": "RATE_LIMIT_EXCEEDED", "message": ( # noqa: E501 - f"Too many requests. Limit: " - f"{self.max_requests} per {self.window_seconds}s" + f"Too many requests. Limit: " f"{self.max_requests} per {self.window_seconds}s" ), } }, @@ -145,8 +139,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 diff --git a/apps/api/app/middleware/prometheus_metrics.py b/apps/api/app/middleware/prometheus_metrics.py index 7b24568c1..a571cee44 100644 --- a/apps/api/app/middleware/prometheus_metrics.py +++ b/apps/api/app/middleware/prometheus_metrics.py @@ -24,7 +24,6 @@ from prometheus_client import ( from starlette.middleware.base import BaseHTTPMiddleware from starlette.responses import PlainTextResponse - # Buckets for HTTP request duration (seconds) HTTP_DURATION_BUCKETS = (0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0) @@ -124,6 +123,7 @@ class PrometheusMetricsMiddleware(BaseHTTPMiddleware): async def metrics_endpoint(request: Request) -> PlainTextResponse: """FastAPI endpoint that returns Prometheus metrics in text format.""" import os + version = os.getenv("APP_VERSION", "unknown") environment = os.getenv("APP_ENV", "unknown") APP_INFO.labels(version=version, environment=environment).set(1) diff --git a/apps/api/app/schemas/chunked_upload.py b/apps/api/app/schemas/chunked_upload.py index cb0d7d122..8c5cbef78 100644 --- a/apps/api/app/schemas/chunked_upload.py +++ b/apps/api/app/schemas/chunked_upload.py @@ -7,7 +7,9 @@ class ChunkedUploadInitRequest(BaseModel): filename: str = Field(..., min_length=1, max_length=255, description="Filename") file_size: int = Field(..., gt=0, le=2147483648, description="File size in bytes, max 2GB") total_chunks: int = Field(..., gt=0, description="Total number of chunks") - content_type: str = Field(default="application/octet-stream", min_length=1, max_length=100, description="Content type") + content_type: str = Field( + default="application/octet-stream", min_length=1, max_length=100, description="Content type" + ) project_id: str = Field(..., min_length=1, description="Project ID") library_id: str = Field(..., min_length=1, description="Asset library ID") diff --git a/apps/api/app/schemas/dashboard.py b/apps/api/app/schemas/dashboard.py index d1732b0a5..bc823c177 100644 --- a/apps/api/app/schemas/dashboard.py +++ b/apps/api/app/schemas/dashboard.py @@ -14,12 +14,14 @@ class RecentTaskItem(BaseModel): class SubscriptionInfo(BaseModel): """用户订阅信息。""" + plan: str = "free" is_active: bool = False class DashboardOverviewResponse(BaseModel): """Dashboard 概览数据。""" + total_assets: int = 0 used_storage_bytes: int = 0 total_titles: int = 0 diff --git a/apps/api/app/schemas/duplication.py b/apps/api/app/schemas/duplication.py index 48416cfbe..eb422e718 100644 --- a/apps/api/app/schemas/duplication.py +++ b/apps/api/app/schemas/duplication.py @@ -1,4 +1,5 @@ """查重 API Pydantic schemas。""" + from __future__ import annotations from pydantic import BaseModel, Field diff --git a/apps/api/app/schemas/generation_task.py b/apps/api/app/schemas/generation_task.py index a9e0beb91..64219f214 100644 --- a/apps/api/app/schemas/generation_task.py +++ b/apps/api/app/schemas/generation_task.py @@ -8,6 +8,7 @@ class CreateGenerationTaskRequest(BaseModel): - 项目模式:project_id + asset_library_id(向后兼容) - 模板模式:template_id + asset_ids / title_ids / voice_ids """ + project_id: str = "" asset_library_id: str = "" strategy_id: str = "" @@ -50,4 +51,5 @@ class GenerationTaskResponse(BaseModel): class ListGenerationTasksResponse(BaseModel): """用户级生成任务列表响应(跨 project)。""" + items: list[GenerationTaskResponse] diff --git a/apps/api/app/schemas/recipe.py b/apps/api/app/schemas/recipe.py index 4dfddec5b..7e75d861e 100644 --- a/apps/api/app/schemas/recipe.py +++ b/apps/api/app/schemas/recipe.py @@ -1,4 +1,5 @@ """Recipe API schemas.""" + from __future__ import annotations from datetime import datetime @@ -6,9 +7,9 @@ from typing import Any, Dict, List, Optional from pydantic import BaseModel, Field - # ── Response ── + class RecipeItemResponse(BaseModel): id: str recipe_id: str @@ -50,6 +51,7 @@ class UseRecipeResponse(BaseModel): # ── Request ── + class RecipeItemRequest(BaseModel): item_type: str item_id: str diff --git a/apps/api/app/schemas/subscription.py b/apps/api/app/schemas/subscription.py index 8c42d0939..2080095fc 100644 --- a/apps/api/app/schemas/subscription.py +++ b/apps/api/app/schemas/subscription.py @@ -1,15 +1,17 @@ """Subscription schemas for API request/response models.""" + from __future__ import annotations from typing import List, Optional from pydantic import BaseModel, Field - # ============ Enums / Types ============ + class PlanType(str): """套餐类型""" + FREE = "free" STANDARD = "standard" PRO = "pro" @@ -18,6 +20,7 @@ class PlanType(str): class SubscriptionStatus(str): """订阅状态""" + ACTIVE = "active" EXPIRED = "expired" CANCELLED = "cancelled" @@ -26,6 +29,7 @@ class SubscriptionStatus(str): class BillingStatus(str): """账单状态""" + PAID = "paid" PENDING = "pending" FAILED = "failed" @@ -34,14 +38,17 @@ class BillingStatus(str): class BillingCycle(str): """计费周期""" + MONTHLY = "monthly" YEARLY = "yearly" # ============ Response Schemas ============ + class SubscriptionInfo(BaseModel): """当前订阅信息""" + id: str plan_id: str plan_name: str @@ -56,6 +63,7 @@ class SubscriptionInfo(BaseModel): class BillingRecord(BaseModel): """账单记录""" + id: str plan_name: str amount: float @@ -68,6 +76,7 @@ class BillingRecord(BaseModel): class ChangePlanResponse(BaseModel): """升级/降级响应""" + success: bool message: str new_subscription: Optional[SubscriptionInfo] = None @@ -75,18 +84,22 @@ class ChangePlanResponse(BaseModel): class SimpleResponse(BaseModel): """简单响应(用于取消订阅、切换自动续费等)""" + success: bool message: str # ============ Request Schemas ============ + class ChangePlanRequest(BaseModel): """升级/降级请求""" + target_plan_id: str = Field(..., description="目标套餐ID") billing_cycle: str = Field(..., description="计费周期: monthly/yearly") class ToggleAutoRenewRequest(BaseModel): """切换自动续费请求""" + enabled: bool = Field(..., description="是否开启自动续费") diff --git a/apps/api/app/schemas/task_center.py b/apps/api/app/schemas/task_center.py index 0aaaf7170..931fae17e 100644 --- a/apps/api/app/schemas/task_center.py +++ b/apps/api/app/schemas/task_center.py @@ -25,6 +25,7 @@ class ListProjectTasksResponse(BaseModel): class UserTaskResponse(BaseModel): """用户级任务响应(跨 project,用于模板模式)。""" + id: str task_type: str project_id: str = "" @@ -42,4 +43,5 @@ class UserTaskResponse(BaseModel): class ListTasksResponse(BaseModel): """用户级任务列表响应(GET /api/v1/tasks)。""" + items: list[UserTaskResponse] = Field(default_factory=list) diff --git a/apps/api/app/schemas/template.py b/apps/api/app/schemas/template.py index 0d7bea514..9d95ecbb4 100644 --- a/apps/api/app/schemas/template.py +++ b/apps/api/app/schemas/template.py @@ -1,4 +1,5 @@ """Template API schemas.""" + from __future__ import annotations from datetime import datetime @@ -6,9 +7,9 @@ from typing import Any, Dict, List, Optional from pydantic import BaseModel, Field - # ── Segment ── + class SegmentResponse(BaseModel): id: str template_id: str @@ -29,6 +30,7 @@ class SegmentRequest(BaseModel): # ── Template Response ── + class TemplateResponse(BaseModel): id: str user_id: str @@ -53,6 +55,7 @@ class ListTemplatesResponse(BaseModel): # ── Template Request ── + class CreateTemplateRequest(BaseModel): name: str mode: str @@ -79,6 +82,7 @@ class UpdateTemplateRequest(BaseModel): # ── Validate ── + class ValidateTemplateRequest(BaseModel): voiceover_duration: Optional[float] = None # 配音实际时长(秒) @@ -96,6 +100,7 @@ class ValidateTemplateResponse(BaseModel): # ── Category ── + class CategoryResponse(BaseModel): id: str user_id: str diff --git a/apps/api/app/schemas/title_library.py b/apps/api/app/schemas/title_library.py index 8af89af77..b955d1169 100644 --- a/apps/api/app/schemas/title_library.py +++ b/apps/api/app/schemas/title_library.py @@ -1,4 +1,5 @@ """Title library Pydantic schemas.""" + from __future__ import annotations from datetime import datetime diff --git a/apps/api/app/schemas/upload.py b/apps/api/app/schemas/upload.py index 97d442a01..90a7a026d 100644 --- a/apps/api/app/schemas/upload.py +++ b/apps/api/app/schemas/upload.py @@ -1,12 +1,13 @@ from pydantic import BaseModel, Field - class UploadAssetRequest(BaseModel): """素材上传请求(multipart form)""" + project_id: str = Field(..., min_length=1, description="项目 ID") library_id: str = Field(..., min_length=1, description="素材库 ID") + class UploadAssetResponse(BaseModel): storage_key: str ingest_job_id: str diff --git a/apps/api/app/schemas/voice_library.py b/apps/api/app/schemas/voice_library.py index f906e0554..b21ea4362 100644 --- a/apps/api/app/schemas/voice_library.py +++ b/apps/api/app/schemas/voice_library.py @@ -1,4 +1,5 @@ """Voice library Pydantic schemas.""" + from __future__ import annotations from datetime import datetime diff --git a/apps/worker/video_processing/dedup.py b/apps/worker/video_processing/dedup.py index 958a71752..852a3c764 100644 --- a/apps/worker/video_processing/dedup.py +++ b/apps/worker/video_processing/dedup.py @@ -1,4 +1,5 @@ """Video deduplication module - compute fingerprints and detect duplicates.""" + import hashlib import json import logging @@ -59,6 +60,7 @@ def compute_color_histogram(image: np.ndarray, bins: int = 32) -> list[float]: @dataclass class VideoFingerprint: """Video fingerprint containing multiple similarity metrics.""" + md5: str keyframe_phashes: list[str] color_histograms: list[list[float]] @@ -66,7 +68,13 @@ class VideoFingerprint: resolution: tuple[int, int] def to_dict(self) -> dict: - return {"md5": self.md5, "keyframe_phashes": self.keyframe_phashes, "color_histograms": self.color_histograms, "duration": self.duration, "resolution": list(self.resolution)} + return { + "md5": self.md5, + "keyframe_phashes": self.keyframe_phashes, + "color_histograms": self.color_histograms, + "duration": self.duration, + "resolution": list(self.resolution), + } class VideoDeduplicator: @@ -80,48 +88,54 @@ class VideoDeduplicator: cap = cv2.VideoCapture(video_path) if not cap.isOpened(): raise RuntimeError(f"Cannot open video: {video_path}") - + fps = cap.get(cv2.CAP_PROP_FPS) frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) duration = frame_count / fps if fps > 0 else 0 width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) - + md5_hash = hashlib.md5() keyframe_phashes = [] color_histograms = [] - + frame_interval = max(1, frame_count // 10) for i in range(0, frame_count, frame_interval): cap.set(cv2.CAP_PROP_POS_FRAMES, i) ret, frame = cap.read() if not ret: continue - + _, buffer = cv2.imencode(".jpg", frame) md5_hash.update(buffer) - + keyframe_phashes.append(compute_phash(frame)) color_histograms.append(compute_color_histogram(frame)) - + cap.release() - - return VideoFingerprint(md5=md5_hash.hexdigest(), keyframe_phashes=keyframe_phashes, color_histograms=color_histograms, duration=duration, resolution=(width, height)) + + return VideoFingerprint( + md5=md5_hash.hexdigest(), + keyframe_phashes=keyframe_phashes, + color_histograms=color_histograms, + duration=duration, + resolution=(width, height), + ) def check_duplicate(self, fingerprint: VideoFingerprint, project_id: str, session: Session) -> Optional[dict]: """Check if video is duplicate of existing one. Returns duplicate info if found.""" video_repo = SQLAlchemyGeneratedVideoRepository(session) existing_videos = video_repo.list_by_project(project_id) - + for existing in existing_videos: if not existing.video_fingerprint: continue - + ef = existing.video_fingerprint - + if fingerprint.md5 == ef.get("md5"): return {"duplicate": True, "duplicate_of": existing.id, "reason": "exact_md5_match", "similarity": 1.0} - + existing_phashes = ef.get("keyframe_phashes", []) if existing_phashes: total_distance = 0 @@ -130,10 +144,15 @@ class VideoDeduplicator: distances = [hamming_distance(phash, ep) for ep in existing_phashes] min_distances.append(min(distances)) avg_distance = sum(min_distances) / len(min_distances) if min_distances else 100 - + if avg_distance < self.PHASH_THRESHOLD: - return {"duplicate": True, "duplicate_of": existing.id, "reason": "phash_similar", "similarity": 1.0 - (avg_distance / 64)} - + return { + "duplicate": True, + "duplicate_of": existing.id, + "reason": "phash_similar", + "similarity": 1.0 - (avg_distance / 64), + } + return None @@ -142,24 +161,26 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict: """Celery task to check if generated video is a duplicate.""" session = SessionLocal() temp_dir = tempfile.mkdtemp() - + try: video_repo = SQLAlchemyGeneratedVideoRepository(session) storage_service = get_storage_service() deduplicator = VideoDeduplicator() - + video = video_repo.get(generated_video_id) if video is None: raise ValueError(f"Generated video {generated_video_id} not found") - + local_path = os.path.join(temp_dir, f"{generated_video_id}.mp4") storage_key = video.file_url.split("/")[-1] - storage_service.download_file(f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4", local_path) - + storage_service.download_file( + f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4", local_path + ) + fingerprint = deduplicator.compute_fingerprint(local_path) - + duplicate_result = deduplicator.check_duplicate(fingerprint, video.project_id, session) - + video.video_fingerprint = fingerprint.to_dict() if duplicate_result: video.is_duplicate = True @@ -167,13 +188,19 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict: else: video.is_duplicate = False video.duplicate_of = None - + video_repo.update(video) session.commit() - + logger.info(f"Duplicate check completed for video {generated_video_id}: is_duplicate={video.is_duplicate}") - - return {"ok": True, "video_id": generated_video_id, "is_duplicate": video.is_duplicate, "duplicate_of": video.duplicate_of, "fingerprint": fingerprint.to_dict()} + + return { + "ok": True, + "video_id": generated_video_id, + "is_duplicate": video.is_duplicate, + "duplicate_of": video.duplicate_of, + "fingerprint": fingerprint.to_dict(), + } except Exception as e: logger.error(f"Duplicate check failed for {generated_video_id}: {str(e)}") session.rollback() @@ -181,4 +208,5 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict: finally: session.close() import shutil + shutil.rmtree(temp_dir, ignore_errors=True) diff --git a/apps/worker/video_processing/editing_modes.py b/apps/worker/video_processing/editing_modes.py index 3b3968a14..3cc3eb403 100644 --- a/apps/worker/video_processing/editing_modes.py +++ b/apps/worker/video_processing/editing_modes.py @@ -140,14 +140,24 @@ class EditingModeProcessor: try: result = subprocess.run( [ - self._ffprobe_bin, "-v", "error", - "-show_entries", "stream=width,height,r_frame_rate,duration,codec_name", - "-show_entries", "format=duration,size", - "-of", "json", video_path, + self._ffprobe_bin, + "-v", + "error", + "-show_entries", + "stream=width,height,r_frame_rate,duration,codec_name", + "-show_entries", + "format=duration,size", + "-of", + "json", + video_path, ], - check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, + check=True, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, ) import json + data = json.loads(result.stdout) streams = data.get("streams", [{}]) video_stream = next((s for s in streams if s.get("codec_type") == "video"), streams[0] if streams else {}) @@ -169,7 +179,9 @@ class EditingModeProcessor: logger.warning(f"Failed to get video info for {video_path}: {e}") return {"width": 0, "height": 0, "fps": 25, "duration": 0, "codec": "unknown", "size": 0} - def _get_pip_position_offset(self, main_width: int, main_height: int, pip_width: int, pip_height: int) -> tuple[int, int]: + def _get_pip_position_offset( + self, main_width: int, main_height: int, pip_width: int, pip_height: int + ) -> tuple[int, int]: """获取画中画位置偏移量""" margin = 10 position_offsets = { @@ -183,16 +195,28 @@ class EditingModeProcessor: def _normalize_video(self, input_path: str, output_path: str) -> dict: """标准化视频格式:先统一帧率,再缩放/填充""" command = [ - self._ffmpeg_bin, "-y", "-i", input_path, - "-r", str(self.config.output_fps), # 先统一帧率 - "-vf", f"scale={self.config.output_width}:{self.config.output_height}:force_original_aspect_ratio=decrease,pad={self.config.output_width}:{self.config.output_height}:(ow-iw)/2:(oh-ih)/2,setsar=1", - "-r", str(self.config.output_fps), - "-c:v", self.config.output_codec, - "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), - "-pix_fmt", "yuv420p", - "-movflags", "+faststart", - "-an", output_path, + self._ffmpeg_bin, + "-y", + "-i", + input_path, + "-r", + str(self.config.output_fps), # 先统一帧率 + "-vf", + f"scale={self.config.output_width}:{self.config.output_height}:force_original_aspect_ratio=decrease,pad={self.config.output_width}:{self.config.output_height}:(ow-iw)/2:(oh-ih)/2,setsar=1", + "-r", + str(self.config.output_fps), + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + "-movflags", + "+faststart", + "-an", + output_path, ] self._run_ffmpeg(command) return self._get_video_info(output_path) @@ -231,11 +255,25 @@ class EditingModeProcessor: offset1 = durations[0] - transition / 2 command = [ - self._ffmpeg_bin, "-y", "-i", normalized_paths[0], "-i", normalized_paths[1], - "-filter_complex", f"[0:v][1:v]xfade=transition=fade:duration={transition}:offset={offset1}[v]", - "-map", "[v]", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", output_path, + self._ffmpeg_bin, + "-y", + "-i", + normalized_paths[0], + "-i", + normalized_paths[1], + "-filter_complex", + f"[0:v][1:v]xfade=transition=fade:duration={transition}:offset={offset1}[v]", + "-map", + "[v]", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + output_path, ] self._run_ffmpeg(command) return output_path @@ -250,8 +288,17 @@ class EditingModeProcessor: f.write(f"file '{os.path.abspath(path)}'\n") command = [ - self._ffmpeg_bin, "-y", "-f", "concat", "-safe", "0", - "-i", concat_file, "-c", "copy", output_path, + self._ffmpeg_bin, + "-y", + "-f", + "concat", + "-safe", + "0", + "-i", + concat_file, + "-c", + "copy", + output_path, ] self._run_ffmpeg(command) @@ -277,7 +324,9 @@ class EditingModeProcessor: pip_width = int(self.config.output_width * self.config.pip_scale) pip_height = int(self.config.output_height * self.config.pip_scale) - x_offset, y_offset = self._get_pip_position_offset(self.config.output_width, self.config.output_height, pip_width, pip_height) + x_offset, y_offset = self._get_pip_position_offset( + self.config.output_width, self.config.output_height, pip_width, pip_height + ) pip_normalized = os.path.join(self.work_dir, f"pip_{os.getpid()}.mp4") pip_info = self._get_video_info(video_paths[1]) @@ -285,19 +334,43 @@ class EditingModeProcessor: if pip_info["duration"] > main_info["duration"]: temp_pip = os.path.join(self.work_dir, f"pip_temp_{os.getpid()}.mp4") command = [ - self._ffmpeg_bin, "-y", "-i", video_paths[1], "-t", str(main_info["duration"]), - "-vf", f"scale={pip_width}:{pip_height}", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", temp_pip, + self._ffmpeg_bin, + "-y", + "-i", + video_paths[1], + "-t", + str(main_info["duration"]), + "-vf", + f"scale={pip_width}:{pip_height}", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + temp_pip, ] self._run_ffmpeg(command) pip_normalized_input = temp_pip else: command = [ - self._ffmpeg_bin, "-y", "-i", video_paths[1], - "-vf", f"scale={pip_width}:{pip_height}", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", pip_normalized, + self._ffmpeg_bin, + "-y", + "-i", + video_paths[1], + "-vf", + f"scale={pip_width}:{pip_height}", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + pip_normalized, ] self._run_ffmpeg(command) pip_normalized_input = pip_normalized @@ -305,21 +378,49 @@ class EditingModeProcessor: if main_info["duration"] > pip_info["duration"]: looped_pip = os.path.join(self.work_dir, f"pip_looped_{os.getpid()}.mp4") command = [ - self._ffmpeg_bin, "-y", "-stream_loop", "-1", "-i", pip_normalized_input, - "-t", str(main_info["duration"]), - "-vf", f"scale={pip_width}:{pip_height}", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", looped_pip, + self._ffmpeg_bin, + "-y", + "-stream_loop", + "-1", + "-i", + pip_normalized_input, + "-t", + str(main_info["duration"]), + "-vf", + f"scale={pip_width}:{pip_height}", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + looped_pip, ] self._run_ffmpeg(command) pip_normalized_input = looped_pip command = [ - self._ffmpeg_bin, "-y", "-i", main_normalized, "-i", pip_normalized_input, - "-filter_complex", f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]", - "-map", "[v]", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", output_path, + self._ffmpeg_bin, + "-y", + "-i", + main_normalized, + "-i", + pip_normalized_input, + "-filter_complex", + f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]", + "-map", + "[v]", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + output_path, ] self._run_ffmpeg(command) @@ -349,38 +450,87 @@ class EditingModeProcessor: if bg_info["duration"] < audio_duration: looped_bg = os.path.join(self.work_dir, f"bg_looped_{os.getpid()}.mp4") command = [ - self._ffmpeg_bin, "-y", "-stream_loop", "-1", "-i", bg_normalized, - "-t", str(audio_duration), - "-vf", f"scale={self.config.output_width}:{self.config.output_height}", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", looped_bg, + self._ffmpeg_bin, + "-y", + "-stream_loop", + "-1", + "-i", + bg_normalized, + "-t", + str(audio_duration), + "-vf", + f"scale={self.config.output_width}:{self.config.output_height}", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + looped_bg, ] self._run_ffmpeg(command) bg_normalized = looped_bg elif bg_info["duration"] > audio_duration: temp_bg = os.path.join(self.work_dir, f"bg_trimmed_{os.getpid()}.mp4") command = [ - self._ffmpeg_bin, "-y", "-i", bg_normalized, "-t", str(audio_duration), - "-c:v", "copy", temp_bg, + self._ffmpeg_bin, + "-y", + "-i", + bg_normalized, + "-t", + str(audio_duration), + "-c:v", + "copy", + temp_bg, ] self._run_ffmpeg(command) bg_normalized = temp_bg blurred_bg = os.path.join(self.work_dir, f"bg_blurred_{os.getpid()}.mp4") command = [ - self._ffmpeg_bin, "-y", "-i", bg_normalized, - "-vf", f"boxblur=5:5,scale={self.config.output_width}:{self.config.output_height}", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", blurred_bg, + self._ffmpeg_bin, + "-y", + "-i", + bg_normalized, + "-vf", + f"boxblur=5:5,scale={self.config.output_width}:{self.config.output_height}", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + blurred_bg, ] self._run_ffmpeg(command) command = [ - self._ffmpeg_bin, "-y", "-i", blurred_bg, "-i", audio_path, - "-filter_complex", "[0:v]drawbox=x=0:y=0:w=iw:h=ih:color=black@0.3:t=fill[v]", - "-map", "[v]", "-map", "1:a", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", "-shortest", output_path, + self._ffmpeg_bin, + "-y", + "-i", + blurred_bg, + "-i", + audio_path, + "-filter_complex", + "[0:v]drawbox=x=0:y=0:w=iw:h=ih:color=black@0.3:t=fill[v]", + "-map", + "[v]", + "-map", + "1:a", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + "-shortest", + output_path, ] self._run_ffmpeg(command) @@ -414,39 +564,97 @@ class EditingModeProcessor: pip_width = int(self.config.output_width * self.config.pip_scale) pip_height = int(self.config.output_height * self.config.pip_scale) - x_offset, y_offset = self._get_pip_position_offset(self.config.output_width, self.config.output_height, pip_width, pip_height) + x_offset, y_offset = self._get_pip_position_offset( + self.config.output_width, self.config.output_height, pip_width, pip_height + ) voice_adjusted = os.path.join(self.work_dir, f"voice_adj_{os.getpid()}.mp4") command = [ - self._ffmpeg_bin, "-y", "-i", voice_normalized, "-t", str(final_duration), - "-vf", f"scale={pip_width}:{pip_height}", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", voice_adjusted, + self._ffmpeg_bin, + "-y", + "-i", + voice_normalized, + "-t", + str(final_duration), + "-vf", + f"scale={pip_width}:{pip_height}", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + voice_adjusted, ] self._run_ffmpeg(command) bg_adjusted = os.path.join(self.work_dir, f"bg_adj_{os.getpid()}.mp4") command = [ - self._ffmpeg_bin, "-y", "-i", bg_normalized, "-t", str(final_duration), - "-c:v", "copy", bg_adjusted, + self._ffmpeg_bin, + "-y", + "-i", + bg_normalized, + "-t", + str(final_duration), + "-c:v", + "copy", + bg_adjusted, ] self._run_ffmpeg(command) if audio_path: command = [ - self._ffmpeg_bin, "-y", "-i", bg_adjusted, "-i", voice_adjusted, "-i", audio_path, - "-filter_complex", f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]", - "-map", "[v]", "-map", "2:a", "-shortest", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", output_path, + self._ffmpeg_bin, + "-y", + "-i", + bg_adjusted, + "-i", + voice_adjusted, + "-i", + audio_path, + "-filter_complex", + f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]", + "-map", + "[v]", + "-map", + "2:a", + "-shortest", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + output_path, ] else: command = [ - self._ffmpeg_bin, "-y", "-i", bg_adjusted, "-i", voice_adjusted, - "-filter_complex", f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]", - "-map", "[v]", "-map", "1:a", "-shortest", - "-c:v", self.config.output_codec, "-preset", self.config.output_preset, - "-crf", str(self.config.output_crf), "-pix_fmt", "yuv420p", output_path, + self._ffmpeg_bin, + "-y", + "-i", + bg_adjusted, + "-i", + voice_adjusted, + "-filter_complex", + f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]", + "-map", + "[v]", + "-map", + "1:a", + "-shortest", + "-c:v", + self.config.output_codec, + "-preset", + self.config.output_preset, + "-crf", + str(self.config.output_crf), + "-pix_fmt", + "yuv420p", + output_path, ] self._run_ffmpeg(command) diff --git a/apps/worker/worker_app/tasks/__init__.py b/apps/worker/worker_app/tasks/__init__.py index 8440125db..a12198ea0 100644 --- a/apps/worker/worker_app/tasks/__init__.py +++ b/apps/worker/worker_app/tasks/__init__.py @@ -7,23 +7,36 @@ def __getattr__(name: str): """Lazy import task functions to avoid loading Celery at module import time.""" if name == "classify_asset": from .classification import classify_asset + return classify_asset elif name == "generate_video": from .generation import generate_video + return generate_video elif name == "healthcheck": from .health import healthcheck + return healthcheck elif name == "ingest_asset": from .ingest import ingest_asset + return ingest_asset elif name == "extract_voice_task": from .voice_extraction import extract_voice_task + return extract_voice_task elif name == "extract_background_task": from .voice_extraction import extract_background_task + return extract_background_task raise AttributeError(f"module {__name__!r} has no attribute {name!r}") -__all__ = ["classify_asset", "generate_video", "healthcheck", "ingest_asset", "extract_voice_task", "extract_background_task"] +__all__ = [ + "classify_asset", + "generate_video", + "healthcheck", + "ingest_asset", + "extract_voice_task", + "extract_background_task", +] diff --git a/apps/worker/worker_app/tasks/asset_analyzer.py b/apps/worker/worker_app/tasks/asset_analyzer.py index 71b1dbbcf..151fbcf32 100755 --- a/apps/worker/worker_app/tasks/asset_analyzer.py +++ b/apps/worker/worker_app/tasks/asset_analyzer.py @@ -27,6 +27,7 @@ logger = logging.getLogger(__name__) @dataclass class VideoInfo: """视频基本信息""" + width: int = 0 height: int = 0 fps: float = 0.0 @@ -40,6 +41,7 @@ class VideoInfo: @dataclass class ColorAnalysis: """色彩分析结果""" + dominant_hue: float = 0.0 # 主色调 (0-360) green_ratio: float = 0.0 # 绿色占比 warm_ratio: float = 0.0 # 暖色调占比 @@ -51,6 +53,7 @@ class ColorAnalysis: @dataclass class MotionAnalysis: """运动分析结果""" + motion_score: float = 0.0 # 运动幅度 (0-1) scene_changes: int = 0 # 场景切换次数 @@ -58,6 +61,7 @@ class MotionAnalysis: @dataclass class AudioAnalysis: """音频分析结果""" + has_audio: bool = False speech_ratio: float = 0.0 # 人声比例 music_ratio: float = 0.0 # 音乐比例 @@ -67,6 +71,7 @@ class AudioAnalysis: @dataclass class ClassificationResult: """分类结果""" + category: AssetClassification confidence: float scores: dict[str, float] = field(default_factory=dict) @@ -75,6 +80,7 @@ class ClassificationResult: @dataclass class QualityScore: """质量评分结果""" + total: float resolution_score: float = 0.0 fps_score: float = 0.0 @@ -86,14 +92,14 @@ class QualityScore: class AssetAnalyzer: """ 轻量级视频素材分析器 - + 使用 FFmpeg + NumPy 进行视频特征分析,不依赖外部 AI API。 """ def __init__(self, video_path: str, temp_dir: str | None = None): """ 初始化分析器 - + Args: video_path: 视频文件路径 temp_dir: 临时目录,用于存储提取的帧 @@ -111,6 +117,7 @@ class AssetAnalyzer: """清理临时目录""" try: import shutil + if os.path.exists(self._temp_dir): shutil.rmtree(self._temp_dir) except Exception: @@ -126,8 +133,10 @@ class AssetAnalyzer: try: cmd = [ "ffprobe", - "-v", "quiet", - "-print_format", "json", + "-v", + "quiet", + "-print_format", + "json", "-show_format", "-show_streams", self.video_path, @@ -174,10 +183,10 @@ class AssetAnalyzer: def extract_frames(self, count: int = 10, max_frames: int = 30) -> list[np.ndarray]: """ 从视频中均匀抽取帧 - + Args: count: 抽取的帧数 - + Returns: 帧数据列表 (RGB 格式) """ @@ -203,11 +212,16 @@ class AssetAnalyzer: cmd = [ "ffmpeg", "-y", # 覆盖输出文件 - "-ss", str(timestamp), - "-i", self.video_path, - "-vframes", "1", - "-q:v", "2", # 高质量 - "-f", "image2", + "-ss", + str(timestamp), + "-i", + self.video_path, + "-vframes", + "1", + "-q:v", + "2", # 高质量 + "-f", + "image2", output_path, ] @@ -234,6 +248,7 @@ class AssetAnalyzer: """加载图片为 numpy 数组 (RGB 格式)""" try: from PIL import Image + img = Image.open(path) if img.mode != "RGB": img = img.convert("RGB") @@ -245,7 +260,7 @@ class AssetAnalyzer: def analyze_color_distribution(self, frames: list[np.ndarray] | None = None) -> ColorAnalysis: """ 分析色彩分布 (HSV 空间) - + Returns: ColorAnalysis 对象 """ @@ -312,7 +327,7 @@ class AssetAnalyzer: def analyze_motion(self, frames: list[np.ndarray] | None = None) -> MotionAnalysis: """ 分析画面运动幅度 - + Returns: MotionAnalysis 对象 """ @@ -353,7 +368,7 @@ class AssetAnalyzer: def analyze_audio(self) -> AudioAnalysis: """ 分析音频特征 - + Returns: AudioAnalysis 对象 """ @@ -371,11 +386,15 @@ class AssetAnalyzer: cmd = [ "ffmpeg", "-y", - "-i", self.video_path, + "-i", + self.video_path, "-vn", # 不要视频 - "-ac", "1", # 单声道 - "-ar", "8000", # 降低采样率 - "-f", "wav", + "-ac", + "1", # 单声道 + "-ar", + "8000", # 降低采样率 + "-f", + "wav", audio_path, ] @@ -389,6 +408,7 @@ class AssetAnalyzer: if result_audio.returncode == 0 and os.path.exists(audio_path): # 读取音频数据 import struct + with open(audio_path, "rb") as f: # 跳过 WAV 头 f.read(44) @@ -396,16 +416,13 @@ class AssetAnalyzer: if len(audio_data) >= 2: # 转换为 numpy 数组 - audio_samples = np.array( - struct.unpack(f"<{len(audio_data)//2}h", audio_data), - dtype=float - ) + audio_samples = np.array(struct.unpack(f"<{len(audio_data)//2}h", audio_data), dtype=float) audio_samples = audio_samples / 32768.0 if len(audio_samples) > 0: # 简单频谱分析 - fft = np.abs(np.fft.rfft(audio_samples[:min(len(audio_samples), 8000)])) - freqs = np.fft.rfftfreq(min(len(audio_samples), 8000), 1/8000) + fft = np.abs(np.fft.rfft(audio_samples[: min(len(audio_samples), 8000)])) + freqs = np.fft.rfftfreq(min(len(audio_samples), 8000), 1 / 8000) # 人声频率: 300-3400 Hz speech_mask = (freqs >= 300) & (freqs <= 3400) @@ -433,7 +450,7 @@ class AssetAnalyzer: def classify(self) -> ClassificationResult: """ 综合分析得出分类结果 - + Returns: ClassificationResult 对象 """ @@ -445,7 +462,7 @@ class AssetAnalyzer: # 计算各类别得分 scores = self._calculate_category_scores(color, motion, audio) - + # 找最高分 if not scores: return ClassificationResult( @@ -472,12 +489,12 @@ class AssetAnalyzer: ) -> dict[str, float]: """ 计算各类别的置信度得分 - + Args: color: 色彩分析结果 motion: 运动分析结果 audio: 音频分析结果 - + Returns: 各类别得分字典 """ @@ -576,7 +593,7 @@ class AssetAnalyzer: def calculate_quality_score(self) -> QualityScore: """ 计算视频质量综合评分 (0-100) - + 评分维度: 1. 分辨率得分 (25分) 2. 帧率得分 (20分) @@ -661,7 +678,7 @@ class AssetAnalyzer: def _score_clarity(self, frames: list[np.ndarray]) -> float: """ 清晰度评分 (满分 20) - + 使用 Laplacian 方差评估画面清晰度 高方差 = 细节丰富 = 高分 """ @@ -679,15 +696,12 @@ class AssetAnalyzer: gray = frame # Laplacian 算子 - laplacian = np.array([ - [0, 1, 0], - [1, -4, 1], - [0, 1, 0] - ], dtype=np.float32) + laplacian = np.array([[0, 1, 0], [1, -4, 1], [0, 1, 0]], dtype=np.float32) # 手动计算卷积 from scipy import signal - laplacian_img = signal.convolve2d(gray.astype(float), laplacian, mode='same') + + laplacian_img = signal.convolve2d(gray.astype(float), laplacian, mode="same") variance = np.var(laplacian_img) variances.append(variance) @@ -706,7 +720,7 @@ class AssetAnalyzer: def _score_stability(self, frames: list[np.ndarray]) -> float: """ 稳定性评分 (满分 15) - + 分析帧间位移方差 画面稳定 = 高分 剧烈抖动 = 低分 @@ -722,14 +736,10 @@ class AssetAnalyzer: scale = 0.25 new_h = int(frames[i].shape[0] * scale) new_w = int(frames[i].shape[1] * scale) - frame1_small = np.array( - Image.fromarray(frames[i]).resize((new_w, new_h)) - ) + frame1_small = np.array(Image.fromarray(frames[i]).resize((new_w, new_h))) new_h2 = int(frames[i + 1].shape[0] * scale) new_w2 = int(frames[i + 1].shape[1] * scale) - frame2_small = np.array( - Image.fromarray(frames[i + 1]).resize((new_w2, new_h2)) - ) + frame2_small = np.array(Image.fromarray(frames[i + 1]).resize((new_w2, new_h2))) # 简单位移检测:灰度差 gray1 = np.mean(frame1_small, axis=2) if len(frame1_small.shape) == 3 else frame1_small @@ -756,10 +766,10 @@ class AssetAnalyzer: def classify_asset_real(video_path: str) -> tuple[str, float]: """ 真实分类入口函数 - + Args: video_path: 视频文件路径 - + Returns: (分类类别, 置信度) """ @@ -775,10 +785,10 @@ def classify_asset_real(video_path: str) -> tuple[str, float]: def calculate_quality_score_real(video_path: str) -> float: """ 质量评分入口函数 - + Args: video_path: 视频文件路径 - + Returns: 质量评分 (0-100) """ diff --git a/apps/worker/worker_app/tasks/classification.py b/apps/worker/worker_app/tasks/classification.py index 091d1d65c..de6460c7f 100755 --- a/apps/worker/worker_app/tasks/classification.py +++ b/apps/worker/worker_app/tasks/classification.py @@ -75,8 +75,7 @@ def classify_asset(self, job_id: str) -> dict: session.commit() logger.info( - f"Classification completed for asset {asset.id}: " - f"category={classification}, confidence={confidence}" + f"Classification completed for asset {asset.id}: " f"category={classification}, confidence={confidence}" ) return { diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 8c8b3c313..1a7b77da5 100755 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -78,12 +78,19 @@ def _probe_duration(local_path: Path) -> float: try: result = subprocess.run( [ - FFPROBE_BIN, "-v", "error", - "-show_entries", "format=duration", - "-of", "default=noprint_wrappers=1:nokey=1", + FFPROBE_BIN, + "-v", + "error", + "-show_entries", + "format=duration", + "-of", + "default=noprint_wrappers=1:nokey=1", str(local_path), ], - check=True, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, + check=True, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, ) # nosec B603 return round(float(result.stdout.strip()), 3) except Exception: @@ -95,10 +102,20 @@ def _create_fallback_clip(output_path: Path, title: str) -> None: safe_title = title.replace(":", "\\:").replace("'", "\\'")[:80] _run_ffmpeg( [ - FFMPEG_BIN, "-y", "-f", "lavfi", - "-i", f"color=c=#111827:s={OUTPUT_WIDTH}x{OUTPUT_HEIGHT}:d={OUTPUT_DURATION_SECONDS}:r={int(OUTPUT_FPS)}", - "-vf", f"drawtext=text='{safe_title}':fontcolor=white:fontsize=48:x=(w-text_w)/2:y=(h-text_h)/2", - "-c:v", "libx264", "-pix_fmt", "yuv420p", "-movflags", "+faststart", + FFMPEG_BIN, + "-y", + "-f", + "lavfi", + "-i", + f"color=c=#111827:s={OUTPUT_WIDTH}x{OUTPUT_HEIGHT}:d={OUTPUT_DURATION_SECONDS}:r={int(OUTPUT_FPS)}", + "-vf", + f"drawtext=text='{safe_title}':fontcolor=white:fontsize=48:x=(w-text_w)/2:y=(h-text_h)/2", + "-c:v", + "libx264", + "-pix_fmt", + "yuv420p", + "-movflags", + "+faststart", str(output_path), ] ) @@ -144,11 +161,16 @@ def _download_library_assets( try: # 查询素材库中的视频素材 - assets = session.query(AssetModel).filter( - AssetModel.asset_library_id == asset_library_id, - AssetModel.status == "ready", - AssetModel.file_type.in_(["video", "video/mp4", "video/quicktime"]), - ).order_by(AssetModel.created_at).all() + assets = ( + session.query(AssetModel) + .filter( + AssetModel.asset_library_id == asset_library_id, + AssetModel.status == "ready", + AssetModel.file_type.in_(["video", "video/mp4", "video/quicktime"]), + ) + .order_by(AssetModel.created_at) + .all() + ) if not assets: logger.info(f"No video assets found in library {asset_library_id}") diff --git a/apps/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py index 7618c2686..096bafea3 100755 --- a/apps/worker/worker_app/tasks/ingest.py +++ b/apps/worker/worker_app/tasks/ingest.py @@ -37,8 +37,10 @@ def extract_media_metadata(file_url: str, media_type: str) -> dict: # 使用 ffprobe 提取视频元数据 cmd = [ "ffprobe", - "-v", "quiet", - "-print_format", "json", + "-v", + "quiet", + "-print_format", + "json", "-show_format", "-show_streams", file_url, diff --git a/apps/worker/worker_app/tasks/voice_extraction.py b/apps/worker/worker_app/tasks/voice_extraction.py index 38641e66b..f31196478 100644 --- a/apps/worker/worker_app/tasks/voice_extraction.py +++ b/apps/worker/worker_app/tasks/voice_extraction.py @@ -1,4 +1,5 @@ """Voice extraction tasks - extract voice tracks and background music from videos.""" + import json import logging import os @@ -31,10 +32,31 @@ class VoiceExtractor: raise RuntimeError(f"FFmpeg failed: {result.stderr}") return result - def extract_voice(self, input_path: str, output_path: str, highpass: int = 200, bandpass_freq: int = 300, bandpass_width: int = 3000, noise_reduction: int = 20) -> str: + def extract_voice( + self, + input_path: str, + output_path: str, + highpass: int = 200, + bandpass_freq: int = 300, + bandpass_width: int = 3000, + noise_reduction: int = 20, + ) -> str: """Extract voice track from video using FFmpeg.""" os.makedirs(os.path.dirname(output_path), exist_ok=True) - cmd = ["ffmpeg", "-y", "-i", input_path, "-af", f"highpass=f={highpass},afftdn=bn={noise_reduction},bandpass=f={bandpass_freq}:width_type=h:width={bandpass_width},loudnorm", "-vn", "-acodec", "libmp3lame", "-q:a", "2", output_path] + cmd = [ + "ffmpeg", + "-y", + "-i", + input_path, + "-af", + f"highpass=f={highpass},afftdn=bn={noise_reduction},bandpass=f={bandpass_freq}:width_type=h:width={bandpass_width},loudnorm", + "-vn", + "-acodec", + "libmp3lame", + "-q:a", + "2", + output_path, + ] self._run_ffmpeg(cmd) logger.info(f"Voice extracted to: {output_path}") return output_path @@ -42,7 +64,20 @@ class VoiceExtractor: def extract_background(self, input_path: str, output_path: str, lowpass: int = 200) -> str: """Extract background music from video.""" os.makedirs(os.path.dirname(output_path), exist_ok=True) - cmd = ["ffmpeg", "-y", "-i", input_path, "-af", f"lowpass=f={lowpass},loudnorm", "-vn", "-acodec", "libmp3lame", "-q:a", "2", output_path] + cmd = [ + "ffmpeg", + "-y", + "-i", + input_path, + "-af", + f"lowpass=f={lowpass},loudnorm", + "-vn", + "-acodec", + "libmp3lame", + "-q:a", + "2", + output_path, + ] self._run_ffmpeg(cmd) logger.info(f"Background extracted to: {output_path}") return output_path @@ -80,6 +115,7 @@ def extract_voice_task(self: Task, asset_id: str) -> dict: finally: session.close() import shutil + shutil.rmtree(temp_dir, ignore_errors=True) @@ -115,4 +151,5 @@ def extract_background_task(self: Task, asset_id: str) -> dict: finally: session.close() import shutil + shutil.rmtree(temp_dir, ignore_errors=True) diff --git a/fix_script.py b/fix_script.py index 38ac73728..34990d410 100644 --- a/fix_script.py +++ b/fix_script.py @@ -1,8 +1,8 @@ -content = open('F:/openclaw-saas/scripts/init_tracker_data.py', 'r', encoding='utf-8').read() +content = open("F:/openclaw-saas/scripts/init_tracker_data.py", "r", encoding="utf-8").read() content = content.replace('"title":', '"name":') content = content.replace('"URGENT"', '"urgent"') content = content.replace('"HIGH"', '"high"') content = content.replace('"MEDIUM"', '"medium"') content = content.replace('"LOW"', '"low"') -open('F:/openclaw-saas/scripts/init_tracker_data.py', 'w', encoding='utf-8').write(content) -print('Fixed all fields') +open("F:/openclaw-saas/scripts/init_tracker_data.py", "w", encoding="utf-8").write(content) +print("Fixed all fields") diff --git a/fix_tracker_encoding.py b/fix_tracker_encoding.py index 947b1d392..544e11879 100644 --- a/fix_tracker_encoding.py +++ b/fix_tracker_encoding.py @@ -4,15 +4,15 @@ import sqlite3 from datetime import datetime # 删除旧数据库,重新创建 -conn = sqlite3.connect('/app/tracker.db') +conn = sqlite3.connect("/app/tracker.db") cursor = conn.cursor() # 删除所有表 -cursor.execute('DROP TABLE IF EXISTS tasks') -cursor.execute('DROP TABLE IF EXISTS milestones') +cursor.execute("DROP TABLE IF EXISTS tasks") +cursor.execute("DROP TABLE IF EXISTS milestones") # 重新创建表 -cursor.execute('''CREATE TABLE tasks ( +cursor.execute("""CREATE TABLE tasks ( id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, description TEXT, @@ -22,9 +22,9 @@ cursor.execute('''CREATE TABLE tasks ( priority TEXT DEFAULT 'medium', created_at TEXT, updated_at TEXT -)''') +)""") -cursor.execute('''CREATE TABLE milestones ( +cursor.execute("""CREATE TABLE milestones ( id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, phase TEXT, @@ -33,119 +33,123 @@ cursor.execute('''CREATE TABLE milestones ( status TEXT DEFAULT 'pending', description TEXT, created_at TEXT -)''') +)""") # Phase 4 已完成的任务(56个) phase4_tasks = [ - ('JWT Service 实现', '实现 access token 和 refresh token', 'completed', 'high'), - ('Password Hasher 实现', 'bcrypt 密码加密 cost=12', 'completed', 'high'), - ('Redis Session Store', '基于 Redis 的 Session 存储', 'completed', 'high'), - ('Email Service 实现', 'SMTP 邮件服务', 'completed', 'high'), - ('用户注册 API', '用户注册接口', 'completed', 'high'), - ('邮箱验证 API', '邮箱验证接口', 'completed', 'high'), - ('用户登录 API', '用户登录接口', 'completed', 'high'), - ('用户登出 API', '用户登出接口', 'completed', 'high'), - ('密码重置 API', '密码重置流程', 'completed', 'medium'), - ('创建工作空间 API', '创建工作空间接口', 'completed', 'high'), - ('邀请成员 API', '邀请成员接口', 'completed', 'high'), - ('接受拒绝邀请 API', '处理邀请接口', 'completed', 'high'), - ('移除成员 API', '移除成员接口', 'completed', 'medium'), - ('离开工作空间 API', '成员离开接口', 'completed', 'medium'), - ('更新成员角色 API', '修改成员角色', 'completed', 'high'), - ('列出工作空间 API', '查询工作空间列表', 'completed', 'medium'), - ('工作空间详情 API', '工作空间详情', 'completed', 'medium'), - ('列出成员 API', '查询成员列表', 'completed', 'medium'), - ('Permission Checker', '权限检查器', 'completed', 'high'), - ('订阅计划定义', 'Free Pro Enterprise', 'completed', 'high'), - ('升级订阅 API', '订阅升级接口', 'completed', 'high'), - ('取消订阅 API', '订阅取消接口', 'completed', 'medium'), - ('配额检查工具', '配额管理工具', 'completed', 'high'), - ('UserRepository 接口', 'User 仓储接口', 'completed', 'high'), - ('UserRepository InMemory 实现', 'InMemory 实现', 'completed', 'high'), - ('WorkspaceRepository 接口', 'Workspace 仓储接口', 'completed', 'high'), - ('WorkspaceRepository InMemory 实现', 'InMemory 实现', 'completed', 'high'), - ('WorkspaceMemberRepository 接口', 'Member 仓储接口', 'completed', 'high'), - ('WorkspaceMemberRepository InMemory 实现', 'InMemory 实现', 'completed', 'high'), - ('WorkspaceInvitationRepository 接口', 'Invitation 仓储接口', 'completed', 'high'), - ('WorkspaceInvitationRepository InMemory 实现', 'InMemory 实现', 'completed', 'high'), - ('SubscriptionRepository 接口', 'Subscription 仓储接口', 'completed', 'high'), - ('SubscriptionRepository InMemory 实现', 'InMemory 实现', 'completed', 'high'), - ('PostgreSQL Repository 实现', 'PostgreSQL 数据库适配器', 'completed', 'high'), - ('Database Migration 脚本', '数据库迁移脚本', 'completed', 'high'), - ('FastAPI 路由层', 'API 路由实现', 'completed', 'high'), - ('API 文档 Swagger', 'Swagger 文档', 'completed', 'medium'), - ('错误处理中间件', '统一错误处理', 'completed', 'high'), - ('参数验证', 'Pydantic 参数验证', 'completed', 'high'), - ('Docker 配置', 'Docker Compose 配置', 'completed', 'high'), - ('Kubernetes 配置', 'K8s 部署配置', 'completed', 'medium'), - ('健康检查接口', 'Health Check API', 'completed', 'high'), - ('Celery Worker 配置', '异步任务配置', 'completed', 'medium'), - ('Redis 缓存集成', 'Redis 缓存', 'completed', 'high'), - ('GitHub Actions CI/CD', 'CI/CD 流水线', 'completed', 'high'), - ('单元测试 170个', '170 个单元测试', 'completed', 'high'), - ('集成测试', '12 个集成测试', 'completed', 'medium'), - ('性能测试', '性能测试用例', 'completed', 'medium'), - ('连接池优化', '5-6x 性能优化', 'completed', 'high'), - ('API 文档编写', 'API 使用文档', 'completed', 'medium'), - ('部署文档', '部署指南', 'completed', 'medium'), - ('开发文档', '开发指南', 'completed', 'medium'), - ('MIT 开源许可', 'MIT License', 'completed', 'low'), - ('README 完善', 'README.md', 'completed', 'medium'), - ('CONTRIBUTING 指南', '贡献指南', 'completed', 'low'), - ('CODE_OF_CONDUCT', '行为准则', 'completed', 'low'), + ("JWT Service 实现", "实现 access token 和 refresh token", "completed", "high"), + ("Password Hasher 实现", "bcrypt 密码加密 cost=12", "completed", "high"), + ("Redis Session Store", "基于 Redis 的 Session 存储", "completed", "high"), + ("Email Service 实现", "SMTP 邮件服务", "completed", "high"), + ("用户注册 API", "用户注册接口", "completed", "high"), + ("邮箱验证 API", "邮箱验证接口", "completed", "high"), + ("用户登录 API", "用户登录接口", "completed", "high"), + ("用户登出 API", "用户登出接口", "completed", "high"), + ("密码重置 API", "密码重置流程", "completed", "medium"), + ("创建工作空间 API", "创建工作空间接口", "completed", "high"), + ("邀请成员 API", "邀请成员接口", "completed", "high"), + ("接受拒绝邀请 API", "处理邀请接口", "completed", "high"), + ("移除成员 API", "移除成员接口", "completed", "medium"), + ("离开工作空间 API", "成员离开接口", "completed", "medium"), + ("更新成员角色 API", "修改成员角色", "completed", "high"), + ("列出工作空间 API", "查询工作空间列表", "completed", "medium"), + ("工作空间详情 API", "工作空间详情", "completed", "medium"), + ("列出成员 API", "查询成员列表", "completed", "medium"), + ("Permission Checker", "权限检查器", "completed", "high"), + ("订阅计划定义", "Free Pro Enterprise", "completed", "high"), + ("升级订阅 API", "订阅升级接口", "completed", "high"), + ("取消订阅 API", "订阅取消接口", "completed", "medium"), + ("配额检查工具", "配额管理工具", "completed", "high"), + ("UserRepository 接口", "User 仓储接口", "completed", "high"), + ("UserRepository InMemory 实现", "InMemory 实现", "completed", "high"), + ("WorkspaceRepository 接口", "Workspace 仓储接口", "completed", "high"), + ("WorkspaceRepository InMemory 实现", "InMemory 实现", "completed", "high"), + ("WorkspaceMemberRepository 接口", "Member 仓储接口", "completed", "high"), + ("WorkspaceMemberRepository InMemory 实现", "InMemory 实现", "completed", "high"), + ("WorkspaceInvitationRepository 接口", "Invitation 仓储接口", "completed", "high"), + ("WorkspaceInvitationRepository InMemory 实现", "InMemory 实现", "completed", "high"), + ("SubscriptionRepository 接口", "Subscription 仓储接口", "completed", "high"), + ("SubscriptionRepository InMemory 实现", "InMemory 实现", "completed", "high"), + ("PostgreSQL Repository 实现", "PostgreSQL 数据库适配器", "completed", "high"), + ("Database Migration 脚本", "数据库迁移脚本", "completed", "high"), + ("FastAPI 路由层", "API 路由实现", "completed", "high"), + ("API 文档 Swagger", "Swagger 文档", "completed", "medium"), + ("错误处理中间件", "统一错误处理", "completed", "high"), + ("参数验证", "Pydantic 参数验证", "completed", "high"), + ("Docker 配置", "Docker Compose 配置", "completed", "high"), + ("Kubernetes 配置", "K8s 部署配置", "completed", "medium"), + ("健康检查接口", "Health Check API", "completed", "high"), + ("Celery Worker 配置", "异步任务配置", "completed", "medium"), + ("Redis 缓存集成", "Redis 缓存", "completed", "high"), + ("GitHub Actions CI/CD", "CI/CD 流水线", "completed", "high"), + ("单元测试 170个", "170 个单元测试", "completed", "high"), + ("集成测试", "12 个集成测试", "completed", "medium"), + ("性能测试", "性能测试用例", "completed", "medium"), + ("连接池优化", "5-6x 性能优化", "completed", "high"), + ("API 文档编写", "API 使用文档", "completed", "medium"), + ("部署文档", "部署指南", "completed", "medium"), + ("开发文档", "开发指南", "completed", "medium"), + ("MIT 开源许可", "MIT License", "completed", "low"), + ("README 完善", "README.md", "completed", "medium"), + ("CONTRIBUTING 指南", "贡献指南", "completed", "low"), + ("CODE_OF_CONDUCT", "行为准则", "completed", "low"), ] # Phase 4 未完成的任务(4个) phase4_pending = [ - ('文件上传 OSS', '阿里云 OSS 文件上传', 'pending', 'medium'), - ('搜索功能', '全文搜索', 'pending', 'medium'), - ('WebSocket 实时通信', 'WebSocket 支持', 'pending', 'low'), - ('Webhook 支持', 'Webhook 事件推送', 'pending', 'low'), + ("文件上传 OSS", "阿里云 OSS 文件上传", "pending", "medium"), + ("搜索功能", "全文搜索", "pending", "medium"), + ("WebSocket 实时通信", "WebSocket 支持", "pending", "low"), + ("Webhook 支持", "Webhook 事件推送", "pending", "low"), ] now = datetime.now().isoformat() # 插入 Phase 4 任务 for name, desc, status, priority in phase4_tasks + phase4_pending: - cursor.execute('''INSERT INTO tasks + cursor.execute( + """INSERT INTO tasks (name, description, status, phase, priority, created_at, updated_at) - VALUES (?, ?, ?, ?, ?, ?, ?)''', - (name, desc, status, 'Phase 4', priority, now, now)) + VALUES (?, ?, ?, ?, ?, ?, ?)""", + (name, desc, status, "Phase 4", priority, now, now), + ) # 插入里程碑 milestones = [ - ('认证与账号体系', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', '用户注册登录密码管理'), - ('多租户权限体系', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', '工作空间成员管理权限控制'), - ('订阅与计费体系', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', '订阅计划配额管理'), - ('Repository 层', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', '数据仓储层实现'), - ('API 层', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', 'FastAPI 接口实现'), - ('测试与部署', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', '测试 Docker CI/CD'), + ("认证与账号体系", "Phase 4", "2026-06-17", "2026-06-17", "completed", "用户注册登录密码管理"), + ("多租户权限体系", "Phase 4", "2026-06-17", "2026-06-17", "completed", "工作空间成员管理权限控制"), + ("订阅与计费体系", "Phase 4", "2026-06-17", "2026-06-17", "completed", "订阅计划配额管理"), + ("Repository 层", "Phase 4", "2026-06-17", "2026-06-17", "completed", "数据仓储层实现"), + ("API 层", "Phase 4", "2026-06-17", "2026-06-17", "completed", "FastAPI 接口实现"), + ("测试与部署", "Phase 4", "2026-06-17", "2026-06-17", "completed", "测试 Docker CI/CD"), ] for name, phase, start, end, status, desc in milestones: - cursor.execute('''INSERT INTO milestones + cursor.execute( + """INSERT INTO milestones (name, phase, start_date, end_date, status, description, created_at) - VALUES (?, ?, ?, ?, ?, ?, ?)''', - (name, phase, start, end, status, desc, now)) + VALUES (?, ?, ?, ?, ?, ?, ?)""", + (name, phase, start, end, status, desc, now), + ) conn.commit() # 验证 cursor.execute('SELECT COUNT(*) FROM tasks WHERE status = "completed"') completed = cursor.fetchone()[0] -cursor.execute('SELECT COUNT(*) FROM tasks') +cursor.execute("SELECT COUNT(*) FROM tasks") total = cursor.fetchone()[0] -print(f'✅ Tracker 修复完成!') -print(f' - 总任务数: {total}') -print(f' - 已完成: {completed}') -print(f' - 待完成: {total - completed}') -print(f' - 完成率: {completed/total*100:.1f}%') +print(f"✅ Tracker 修复完成!") +print(f" - 总任务数: {total}") +print(f" - 已完成: {completed}") +print(f" - 待完成: {total - completed}") +print(f" - 完成率: {completed/total*100:.1f}%") # 测试中文显示 -cursor.execute('SELECT name FROM tasks LIMIT 3') -print(f'\n前3个任务:') +cursor.execute("SELECT name FROM tasks LIMIT 3") +print(f"\n前3个任务:") for row in cursor.fetchall(): - print(f' - {row[0]}') + print(f" - {row[0]}") conn.close() diff --git a/init_tracker_phase4.py b/init_tracker_phase4.py index 66cbf1d7d..aa333f289 100644 --- a/init_tracker_phase4.py +++ b/init_tracker_phase4.py @@ -1,11 +1,11 @@ import sqlite3 from datetime import datetime -conn = sqlite3.connect('/app/tracker.db') +conn = sqlite3.connect("/app/tracker.db") cursor = conn.cursor() # 创建表 -cursor.execute('''CREATE TABLE IF NOT EXISTS tasks ( +cursor.execute("""CREATE TABLE IF NOT EXISTS tasks ( id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, description TEXT, @@ -15,9 +15,9 @@ cursor.execute('''CREATE TABLE IF NOT EXISTS tasks ( priority TEXT DEFAULT 'medium', created_at TEXT, updated_at TEXT -)''') +)""") -cursor.execute('''CREATE TABLE IF NOT EXISTS milestones ( +cursor.execute("""CREATE TABLE IF NOT EXISTS milestones ( id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, phase TEXT, @@ -26,113 +26,117 @@ cursor.execute('''CREATE TABLE IF NOT EXISTS milestones ( status TEXT DEFAULT 'pending', description TEXT, created_at TEXT -)''') +)""") # Phase 4 已完成的任务(56个) phase4_tasks = [ - ('JWT Service 实现', '实现 access token 和 refresh token', 'completed', 'high'), - ('Password Hasher 实现', 'bcrypt 密码加密,cost=12', 'completed', 'high'), - ('Redis Session Store', '基于 Redis 的 Session 存储', 'completed', 'high'), - ('Email Service 实现', 'SMTP 邮件服务', 'completed', 'high'), - ('用户注册 API', '用户注册接口', 'completed', 'high'), - ('邮箱验证 API', '邮箱验证接口', 'completed', 'high'), - ('用户登录 API', '用户登录接口', 'completed', 'high'), - ('用户登出 API', '用户登出接口', 'completed', 'high'), - ('密码重置 API', '密码重置流程', 'completed', 'medium'), - ('创建工作空间 API', '创建工作空间接口', 'completed', 'high'), - ('邀请成员 API', '邀请成员接口', 'completed', 'high'), - ('接受/拒绝邀请 API', '处理邀请接口', 'completed', 'high'), - ('移除成员 API', '移除成员接口', 'completed', 'medium'), - ('离开工作空间 API', '成员离开接口', 'completed', 'medium'), - ('更新成员角色 API', '修改成员角色', 'completed', 'high'), - ('列出工作空间 API', '查询工作空间列表', 'completed', 'medium'), - ('工作空间详情 API', '工作空间详情', 'completed', 'medium'), - ('列出成员 API', '查询成员列表', 'completed', 'medium'), - ('Permission Checker', '权限检查器', 'completed', 'high'), - ('订阅计划定义', 'Free/Pro/Enterprise', 'completed', 'high'), - ('升级订阅 API', '订阅升级接口', 'completed', 'high'), - ('取消订阅 API', '订阅取消接口', 'completed', 'medium'), - ('配额检查工具', '配额管理工具', 'completed', 'high'), - ('UserRepository 接口', 'User 仓储接口', 'completed', 'high'), - ('UserRepository InMemory 实现', 'InMemory 实现', 'completed', 'high'), - ('WorkspaceRepository 接口', 'Workspace 仓储接口', 'completed', 'high'), - ('WorkspaceRepository InMemory 实现', 'InMemory 实现', 'completed', 'high'), - ('WorkspaceMemberRepository 接口', 'Member 仓储接口', 'completed', 'high'), - ('WorkspaceMemberRepository InMemory 实现', 'InMemory 实现', 'completed', 'high'), - ('WorkspaceInvitationRepository 接口', 'Invitation 仓储接口', 'completed', 'high'), - ('WorkspaceInvitationRepository InMemory 实现', 'InMemory 实现', 'completed', 'high'), - ('SubscriptionRepository 接口', 'Subscription 仓储接口', 'completed', 'high'), - ('SubscriptionRepository InMemory 实现', 'InMemory 实现', 'completed', 'high'), - ('PostgreSQL Repository 实现', 'PostgreSQL 数据库适配器', 'completed', 'high'), - ('Database Migration 脚本', '数据库迁移脚本', 'completed', 'high'), - ('FastAPI 路由层', 'API 路由实现', 'completed', 'high'), - ('API 文档(Swagger)', 'Swagger 文档', 'completed', 'medium'), - ('错误处理中间件', '统一错误处理', 'completed', 'high'), - ('参数验证', 'Pydantic 参数验证', 'completed', 'high'), - ('Docker 配置', 'Docker Compose 配置', 'completed', 'high'), - ('Kubernetes 配置', 'K8s 部署配置', 'completed', 'medium'), - ('健康检查接口', 'Health Check API', 'completed', 'high'), - ('Celery Worker 配置', '异步任务配置', 'completed', 'medium'), - ('Redis 缓存集成', 'Redis 缓存', 'completed', 'high'), - ('GitHub Actions CI/CD', 'CI/CD 流水线', 'completed', 'high'), - ('单元测试(170个)', '170 个单元测试', 'completed', 'high'), - ('集成测试', '12 个集成测试', 'completed', 'medium'), - ('性能测试', '性能测试用例', 'completed', 'medium'), - ('连接池优化', '5-6x 性能优化', 'completed', 'high'), - ('API 文档编写', 'API 使用文档', 'completed', 'medium'), - ('部署文档', '部署指南', 'completed', 'medium'), - ('开发文档', '开发指南', 'completed', 'medium'), - ('MIT 开源许可', 'MIT License', 'completed', 'low'), - ('README 完善', 'README.md', 'completed', 'medium'), - ('CONTRIBUTING 指南', '贡献指南', 'completed', 'low'), - ('CODE_OF_CONDUCT', '行为准则', 'completed', 'low'), + ("JWT Service 实现", "实现 access token 和 refresh token", "completed", "high"), + ("Password Hasher 实现", "bcrypt 密码加密,cost=12", "completed", "high"), + ("Redis Session Store", "基于 Redis 的 Session 存储", "completed", "high"), + ("Email Service 实现", "SMTP 邮件服务", "completed", "high"), + ("用户注册 API", "用户注册接口", "completed", "high"), + ("邮箱验证 API", "邮箱验证接口", "completed", "high"), + ("用户登录 API", "用户登录接口", "completed", "high"), + ("用户登出 API", "用户登出接口", "completed", "high"), + ("密码重置 API", "密码重置流程", "completed", "medium"), + ("创建工作空间 API", "创建工作空间接口", "completed", "high"), + ("邀请成员 API", "邀请成员接口", "completed", "high"), + ("接受/拒绝邀请 API", "处理邀请接口", "completed", "high"), + ("移除成员 API", "移除成员接口", "completed", "medium"), + ("离开工作空间 API", "成员离开接口", "completed", "medium"), + ("更新成员角色 API", "修改成员角色", "completed", "high"), + ("列出工作空间 API", "查询工作空间列表", "completed", "medium"), + ("工作空间详情 API", "工作空间详情", "completed", "medium"), + ("列出成员 API", "查询成员列表", "completed", "medium"), + ("Permission Checker", "权限检查器", "completed", "high"), + ("订阅计划定义", "Free/Pro/Enterprise", "completed", "high"), + ("升级订阅 API", "订阅升级接口", "completed", "high"), + ("取消订阅 API", "订阅取消接口", "completed", "medium"), + ("配额检查工具", "配额管理工具", "completed", "high"), + ("UserRepository 接口", "User 仓储接口", "completed", "high"), + ("UserRepository InMemory 实现", "InMemory 实现", "completed", "high"), + ("WorkspaceRepository 接口", "Workspace 仓储接口", "completed", "high"), + ("WorkspaceRepository InMemory 实现", "InMemory 实现", "completed", "high"), + ("WorkspaceMemberRepository 接口", "Member 仓储接口", "completed", "high"), + ("WorkspaceMemberRepository InMemory 实现", "InMemory 实现", "completed", "high"), + ("WorkspaceInvitationRepository 接口", "Invitation 仓储接口", "completed", "high"), + ("WorkspaceInvitationRepository InMemory 实现", "InMemory 实现", "completed", "high"), + ("SubscriptionRepository 接口", "Subscription 仓储接口", "completed", "high"), + ("SubscriptionRepository InMemory 实现", "InMemory 实现", "completed", "high"), + ("PostgreSQL Repository 实现", "PostgreSQL 数据库适配器", "completed", "high"), + ("Database Migration 脚本", "数据库迁移脚本", "completed", "high"), + ("FastAPI 路由层", "API 路由实现", "completed", "high"), + ("API 文档(Swagger)", "Swagger 文档", "completed", "medium"), + ("错误处理中间件", "统一错误处理", "completed", "high"), + ("参数验证", "Pydantic 参数验证", "completed", "high"), + ("Docker 配置", "Docker Compose 配置", "completed", "high"), + ("Kubernetes 配置", "K8s 部署配置", "completed", "medium"), + ("健康检查接口", "Health Check API", "completed", "high"), + ("Celery Worker 配置", "异步任务配置", "completed", "medium"), + ("Redis 缓存集成", "Redis 缓存", "completed", "high"), + ("GitHub Actions CI/CD", "CI/CD 流水线", "completed", "high"), + ("单元测试(170个)", "170 个单元测试", "completed", "high"), + ("集成测试", "12 个集成测试", "completed", "medium"), + ("性能测试", "性能测试用例", "completed", "medium"), + ("连接池优化", "5-6x 性能优化", "completed", "high"), + ("API 文档编写", "API 使用文档", "completed", "medium"), + ("部署文档", "部署指南", "completed", "medium"), + ("开发文档", "开发指南", "completed", "medium"), + ("MIT 开源许可", "MIT License", "completed", "low"), + ("README 完善", "README.md", "completed", "medium"), + ("CONTRIBUTING 指南", "贡献指南", "completed", "low"), + ("CODE_OF_CONDUCT", "行为准则", "completed", "low"), ] # Phase 4 未完成的任务(4个) phase4_pending = [ - ('文件上传(OSS)', '阿里云 OSS 文件上传', 'pending', 'medium'), - ('搜索功能', '全文搜索', 'pending', 'medium'), - ('WebSocket 实时通信', 'WebSocket 支持', 'pending', 'low'), - ('Webhook 支持', 'Webhook 事件推送', 'pending', 'low'), + ("文件上传(OSS)", "阿里云 OSS 文件上传", "pending", "medium"), + ("搜索功能", "全文搜索", "pending", "medium"), + ("WebSocket 实时通信", "WebSocket 支持", "pending", "low"), + ("Webhook 支持", "Webhook 事件推送", "pending", "low"), ] now = datetime.now().isoformat() # 插入 Phase 4 任务 for name, desc, status, priority in phase4_tasks + phase4_pending: - cursor.execute('''INSERT INTO tasks + cursor.execute( + """INSERT INTO tasks (name, description, status, phase, priority, created_at, updated_at) - VALUES (?, ?, ?, ?, ?, ?, ?)''', - (name, desc, status, 'Phase 4', priority, now, now)) + VALUES (?, ?, ?, ?, ?, ?, ?)""", + (name, desc, status, "Phase 4", priority, now, now), + ) # 插入里程碑 milestones = [ - ('认证与账号体系', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', '用户注册、登录、密码管理'), - ('多租户权限体系', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', '工作空间、成员管理、权限控制'), - ('订阅与计费体系', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', '订阅计划、配额管理'), - ('Repository 层', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', '数据仓储层实现'), - ('API 层', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', 'FastAPI 接口实现'), - ('测试与部署', 'Phase 4', '2026-06-17', '2026-06-17', 'completed', '测试、Docker、CI/CD'), + ("认证与账号体系", "Phase 4", "2026-06-17", "2026-06-17", "completed", "用户注册、登录、密码管理"), + ("多租户权限体系", "Phase 4", "2026-06-17", "2026-06-17", "completed", "工作空间、成员管理、权限控制"), + ("订阅与计费体系", "Phase 4", "2026-06-17", "2026-06-17", "completed", "订阅计划、配额管理"), + ("Repository 层", "Phase 4", "2026-06-17", "2026-06-17", "completed", "数据仓储层实现"), + ("API 层", "Phase 4", "2026-06-17", "2026-06-17", "completed", "FastAPI 接口实现"), + ("测试与部署", "Phase 4", "2026-06-17", "2026-06-17", "completed", "测试、Docker、CI/CD"), ] for name, phase, start, end, status, desc in milestones: - cursor.execute('''INSERT INTO milestones + cursor.execute( + """INSERT INTO milestones (name, phase, start_date, end_date, status, description, created_at) - VALUES (?, ?, ?, ?, ?, ?, ?)''', - (name, phase, start, end, status, desc, now)) + VALUES (?, ?, ?, ?, ?, ?, ?)""", + (name, phase, start, end, status, desc, now), + ) conn.commit() # 统计 cursor.execute('SELECT COUNT(*) FROM tasks WHERE status = "completed"') completed = cursor.fetchone()[0] -cursor.execute('SELECT COUNT(*) FROM tasks') +cursor.execute("SELECT COUNT(*) FROM tasks") total = cursor.fetchone()[0] -print(f'✅ Tracker 初始化完成!') -print(f' - 总任务数: {total}') -print(f' - 已完成: {completed}') -print(f' - 待完成: {total - completed}') -print(f' - 完成率: {completed/total*100:.1f}%') +print(f"✅ Tracker 初始化完成!") +print(f" - 总任务数: {total}") +print(f" - 已完成: {completed}") +print(f" - 待完成: {total - completed}") +print(f" - 完成率: {completed/total*100:.1f}%") conn.close() diff --git a/init_tracker_simple.py b/init_tracker_simple.py index d0d193a8c..09f695db8 100644 --- a/init_tracker_simple.py +++ b/init_tracker_simple.py @@ -1,11 +1,11 @@ import sqlite3 from datetime import datetime -conn = sqlite3.connect('/app/tracker.db') +conn = sqlite3.connect("/app/tracker.db") cursor = conn.cursor() # 创建表 -cursor.execute('''CREATE TABLE IF NOT EXISTS tasks ( +cursor.execute("""CREATE TABLE IF NOT EXISTS tasks ( id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, description TEXT, @@ -14,9 +14,9 @@ cursor.execute('''CREATE TABLE IF NOT EXISTS tasks ( milestone TEXT, priority TEXT DEFAULT 'medium', created_at TEXT -)''') +)""") -cursor.execute('''CREATE TABLE IF NOT EXISTS milestones ( +cursor.execute("""CREATE TABLE IF NOT EXISTS milestones ( id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, phase TEXT, @@ -24,17 +24,17 @@ cursor.execute('''CREATE TABLE IF NOT EXISTS milestones ( end_date TEXT, status TEXT DEFAULT 'pending', description TEXT -)''') +)""") -cursor.execute('''CREATE TABLE IF NOT EXISTS logs ( +cursor.execute("""CREATE TABLE IF NOT EXISTS logs ( id INTEGER PRIMARY KEY AUTOINCREMENT, task_id INTEGER, message TEXT, created_at TEXT -)''') +)""") conn.commit() -print('[OK] Database structure created') +print("[OK] Database structure created") # 插入 Phase 4 和 Phase 6 数据 # Phase 4 里程碑和任务 @@ -52,8 +52,10 @@ milestones = [ ] for name, phase, start, end, status in milestones: - cursor.execute("INSERT INTO milestones (name, phase, start_date, end_date, status) VALUES (?, ?, ?, ?, ?)", - (name, phase, start, end, status)) + cursor.execute( + "INSERT INTO milestones (name, phase, start_date, end_date, status) VALUES (?, ?, ?, ?, ?)", + (name, phase, start, end, status), + ) # Phase 4 任务 (30个) phase4_tasks = [ @@ -90,8 +92,10 @@ phase4_tasks = [ ] for name, milestone, status, priority in phase4_tasks: - cursor.execute("INSERT INTO tasks (name, milestone, status, phase, priority, created_at) VALUES (?, ?, ?, ?, ?, ?)", - (name, milestone, status, "Phase 4", priority, datetime.now().isoformat())) + cursor.execute( + "INSERT INTO tasks (name, milestone, status, phase, priority, created_at) VALUES (?, ?, ?, ?, ?, ?)", + (name, milestone, status, "Phase 4", priority, datetime.now().isoformat()), + ) # Phase 6 任务 (40个) phase6_tasks = [ @@ -138,10 +142,12 @@ phase6_tasks = [ ] for name, milestone, status, priority in phase6_tasks: - cursor.execute("INSERT INTO tasks (name, milestone, status, phase, priority, created_at) VALUES (?, ?, ?, ?, ?, ?)", - (name, milestone, status, "Phase 6", priority, datetime.now().isoformat())) + cursor.execute( + "INSERT INTO tasks (name, milestone, status, phase, priority, created_at) VALUES (?, ?, ?, ?, ?, ?)", + (name, milestone, status, "Phase 6", priority, datetime.now().isoformat()), + ) conn.commit() conn.close() -print('[SUCCESS] tracker.db initialized with 70 tasks and 10 milestones') +print("[SUCCESS] tracker.db initialized with 70 tasks and 10 milestones") diff --git a/packages/adapters/sqlalchemy_impl/asset_library_repository.py b/packages/adapters/sqlalchemy_impl/asset_library_repository.py index 5b418084c..52025d794 100644 --- a/packages/adapters/sqlalchemy_impl/asset_library_repository.py +++ b/packages/adapters/sqlalchemy_impl/asset_library_repository.py @@ -27,9 +27,7 @@ class SQLAlchemyAssetLibraryRepository: return self.get(library_id) def find_by_project(self, project_id: str) -> list[AssetLibrary]: - models = self.session.query(AssetLibraryModel).filter( - AssetLibraryModel.project_id == project_id - ).all() + models = self.session.query(AssetLibraryModel).filter(AssetLibraryModel.project_id == project_id).all() return [ AssetLibrary( id=model.id, @@ -60,9 +58,7 @@ class SQLAlchemyAssetLibraryRepository: return library def update(self, library: AssetLibrary) -> AssetLibrary: - model = self.session.query(AssetLibraryModel).filter( - AssetLibraryModel.id == library.id - ).first() + model = self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library.id).first() if model: model.project_id = library.project_id model.name = library.name @@ -74,9 +70,7 @@ class SQLAlchemyAssetLibraryRepository: return library def delete(self, library_id: str) -> bool: - model = self.session.query(AssetLibraryModel).filter( - AssetLibraryModel.id == library_id - ).first() + model = self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library_id).first() if model: self.session.delete(model) self.session.commit() @@ -84,18 +78,14 @@ class SQLAlchemyAssetLibraryRepository: return False async def increment_asset_count(self, library_id: str, size_delta: int) -> None: - model = self.session.query(AssetLibraryModel).filter( - AssetLibraryModel.id == library_id - ).first() + model = self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library_id).first() if model: model.asset_count = (model.asset_count or 0) + 1 model.total_size = (model.total_size or 0) + size_delta self.session.commit() async def decrement_asset_count(self, library_id: str, size_delta: int) -> None: - model = self.session.query(AssetLibraryModel).filter( - AssetLibraryModel.id == library_id - ).first() + model = self.session.query(AssetLibraryModel).filter(AssetLibraryModel.id == library_id).first() if model: model.asset_count = max(0, (model.asset_count or 0) - 1) model.total_size = max(0, (model.total_size or 0) - size_delta) diff --git a/packages/adapters/sqlalchemy_impl/asset_repository.py b/packages/adapters/sqlalchemy_impl/asset_repository.py index b53cecf3b..ce1e8c9b7 100644 --- a/packages/adapters/sqlalchemy_impl/asset_repository.py +++ b/packages/adapters/sqlalchemy_impl/asset_repository.py @@ -17,9 +17,13 @@ class SQLAlchemyAssetRepository: skip: int = 0, limit: int = 100, ) -> list[Asset]: - models = self.session.query(AssetModel).filter( - AssetModel.asset_library_id == library_id - ).offset(skip).limit(limit).all() + models = ( + self.session.query(AssetModel) + .filter(AssetModel.asset_library_id == library_id) + .offset(skip) + .limit(limit) + .all() + ) return [self._to_domain(model) for model in models] def find_by_project( @@ -28,9 +32,9 @@ class SQLAlchemyAssetRepository: skip: int = 0, limit: int = 100, ) -> list[Asset]: - models = self.session.query(AssetModel).filter( - AssetModel.project_id == project_id - ).offset(skip).limit(limit).all() + models = ( + self.session.query(AssetModel).filter(AssetModel.project_id == project_id).offset(skip).limit(limit).all() + ) return [self._to_domain(model) for model in models] def find_by_id(self, asset_id: str) -> Asset | None: @@ -101,24 +105,23 @@ class SQLAlchemyAssetRepository: return False def count_by_project(self, project_id: str) -> int: - return self.session.query(AssetModel).filter( - AssetModel.project_id == project_id - ).count() + return self.session.query(AssetModel).filter(AssetModel.project_id == project_id).count() def count_by_project_ids(self, project_ids: list[str]) -> int: if not project_ids: return 0 - return self.session.query(AssetModel).filter( - AssetModel.project_id.in_(project_ids) - ).count() + return self.session.query(AssetModel).filter(AssetModel.project_id.in_(project_ids)).count() def sum_storage_by_project_ids(self, project_ids: list[str]) -> int: if not project_ids: return 0 from sqlalchemy import func - result = self.session.query(func.coalesce(func.sum(AssetModel.file_size), 0)).filter( - AssetModel.project_id.in_(project_ids) - ).scalar() + + result = ( + self.session.query(func.coalesce(func.sum(AssetModel.file_size), 0)) + .filter(AssetModel.project_id.in_(project_ids)) + .scalar() + ) return int(result or 0) def _to_domain(self, model: AssetModel) -> Asset: diff --git a/packages/adapters/sqlalchemy_impl/duplication_repository.py b/packages/adapters/sqlalchemy_impl/duplication_repository.py index 9211d8525..f8806dd90 100644 --- a/packages/adapters/sqlalchemy_impl/duplication_repository.py +++ b/packages/adapters/sqlalchemy_impl/duplication_repository.py @@ -1,4 +1,5 @@ """查重记录 SQLAlchemy 仓库实现。""" + from __future__ import annotations import json @@ -34,9 +35,7 @@ class SQLAlchemyDuplicationRecordRepository: return record def get(self, record_id: str) -> DuplicationRecord | None: - model = self.session.query(DuplicationRecordModel).filter( - DuplicationRecordModel.id == record_id - ).first() + model = self.session.query(DuplicationRecordModel).filter(DuplicationRecordModel.id == record_id).first() if model is None: return None return self._to_domain(model) @@ -53,9 +52,7 @@ class SQLAlchemyDuplicationRecordRepository: return [self._to_domain(m) for m in models] def update(self, record: DuplicationRecord) -> DuplicationRecord: - model = self.session.query(DuplicationRecordModel).filter( - DuplicationRecordModel.id == record.id - ).first() + model = self.session.query(DuplicationRecordModel).filter(DuplicationRecordModel.id == record.id).first() if model is None: return record model.status = record.status @@ -66,9 +63,7 @@ class SQLAlchemyDuplicationRecordRepository: model.updated_at = record.updated_at # 更新 segments:先删后建 - self.session.query(DuplicationSegmentModel).filter( - DuplicationSegmentModel.record_id == record.id - ).delete() + self.session.query(DuplicationSegmentModel).filter(DuplicationSegmentModel.record_id == record.id).delete() for seg in record.segments: seg_model = DuplicationSegmentModel( id=seg.id, @@ -87,20 +82,14 @@ class SQLAlchemyDuplicationRecordRepository: return record def delete(self, record_id: str) -> bool: - count = self.session.query(DuplicationRecordModel).filter( - DuplicationRecordModel.id == record_id - ).delete() - self.session.query(DuplicationSegmentModel).filter( - DuplicationSegmentModel.record_id == record_id - ).delete() + count = self.session.query(DuplicationRecordModel).filter(DuplicationRecordModel.id == record_id).delete() + self.session.query(DuplicationSegmentModel).filter(DuplicationSegmentModel.record_id == record_id).delete() self.session.commit() return count > 0 def _to_domain(self, model: DuplicationRecordModel) -> DuplicationRecord: segment_models = ( - self.session.query(DuplicationSegmentModel) - .filter(DuplicationSegmentModel.record_id == model.id) - .all() + self.session.query(DuplicationSegmentModel).filter(DuplicationSegmentModel.record_id == model.id).all() ) segments = [ DuplicateSegment( diff --git a/packages/adapters/sqlalchemy_impl/generation_task_repository.py b/packages/adapters/sqlalchemy_impl/generation_task_repository.py index 7cbdc0d61..eff5a73fb 100644 --- a/packages/adapters/sqlalchemy_impl/generation_task_repository.py +++ b/packages/adapters/sqlalchemy_impl/generation_task_repository.py @@ -80,11 +80,7 @@ class SQLAlchemyGenerationTaskRepository: return [_to_domain(m) for m in models] def count_by_user(self, user_id: str) -> int: - return ( - self.session.query(GenerationTaskModel) - .filter(GenerationTaskModel.created_by_user_id == user_id) - .count() - ) + return self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id).count() def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: models = ( diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 9a6a487b8..c0a5907ec 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -39,7 +39,7 @@ class ProjectModel(Base): name = Column(String(100), nullable=False) description = Column(Text, nullable=False, default="") shared_users = Column(JSON, nullable=False, default=list) # 被共享的用户 ID 列表 - extra_meta = Column('metadata', JSON, nullable=False, default=dict) + extra_meta = Column("metadata", JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) @@ -52,7 +52,7 @@ class AssetLibraryModel(Base): kind = Column(String(20), nullable=False, index=True) asset_count = Column(Float, nullable=False, default=0) total_size = Column(Float, nullable=False, default=0) - extra_meta = Column('metadata', JSON, nullable=False, default=dict) + extra_meta = Column("metadata", JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) @@ -80,12 +80,11 @@ class AssetModel(Base): classification_result = Column(Text, nullable=True) quality_score = Column(Float, nullable=True) uploaded_by_user_id = Column(String(36), nullable=False) - extra_meta = Column('metadata', JSON, nullable=False, default=dict) + extra_meta = Column("metadata", JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc), index=True) updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) - class EditTemplateModel(Base): __tablename__ = "edit_templates" @@ -97,14 +96,11 @@ class EditTemplateModel(Base): clip_count = Column(Integer, nullable=False, default=3) is_active = Column(Boolean, nullable=False, default=True) created_by_user_id = Column(String(32), nullable=False, default="") - extra_meta = Column('metadata', JSON, nullable=False, default=dict) + extra_meta = Column("metadata", JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) - - - class IngestJobModel(Base): __tablename__ = "ingest_jobs" @@ -145,7 +141,9 @@ class GenerationTaskModel(Base): asset_ids = Column(JSON, nullable=False, default=list) title_ids = Column(JSON, nullable=False, default=list) voice_ids = Column(JSON, nullable=False, default=list) - editing_mode = Column(String(20), nullable=False, default="one_take", index=True) # 剪辑模式: one_take, pip, voice_over, voice_pip + editing_mode = Column( + String(20), nullable=False, default="one_take", index=True + ) # 剪辑模式: one_take, pip, voice_over, voice_pip status = Column(String(20), nullable=False, default="pending", index=True) progress = Column(Float, nullable=False, default=0.0) result_count = Column(Float, nullable=False, default=0) @@ -153,7 +151,7 @@ class GenerationTaskModel(Base): started_at = Column(DateTime, nullable=True) completed_at = Column(DateTime, nullable=True) created_by_user_id = Column(String(32), nullable=False, default="", index=True) - extra_meta = Column('metadata', JSON, nullable=False, default=dict) + extra_meta = Column("metadata", JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) @@ -184,9 +182,6 @@ class GeneratedVideoModel(Base): duplicate_of = Column(String(32), nullable=True) - - - class TitleLibraryModel(Base): __tablename__ = "title_libraries" @@ -199,7 +194,7 @@ class TitleLibraryModel(Base): tags = Column(JSON, nullable=False, default=list) usage_count = Column(Integer, nullable=False, default=0) is_active = Column(Boolean, nullable=False, default=True, index=True) - extra_meta = Column('metadata', JSON, nullable=False, default=dict) + extra_meta = Column("metadata", JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) @@ -220,7 +215,7 @@ class VoiceLibraryModel(Base): file_size = Column(Integer, nullable=False, default=0) status = Column(String(20), nullable=False, default="completed", index=True) tags = Column(JSON, nullable=False, default=list) - extra_meta = Column('metadata', JSON, nullable=False, default=dict) + extra_meta = Column("metadata", JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) @@ -267,7 +262,7 @@ class RecipeModel(Base): template_id = Column(String(36), nullable=False, default="") generation_params = Column(JSON, nullable=False, default=dict) is_active = Column(Boolean, nullable=False, default=True) - extra_meta = Column('metadata', JSON, nullable=False, default=dict) + extra_meta = Column("metadata", JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) @@ -280,7 +275,7 @@ class RecipeItemModel(Base): item_type = Column(String(20), nullable=False) item_id = Column(String(36), nullable=False) position = Column(Integer, nullable=False, default=0) - extra_meta = Column('metadata', JSON, nullable=False, default=dict) + extra_meta = Column("metadata", JSON, nullable=False, default=dict) class TemplateModel(Base): @@ -321,4 +316,3 @@ class TemplateCategoryModel(Base): user_id = Column(String(36), nullable=False, index=True) name = Column(String(100), nullable=False) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) - diff --git a/packages/adapters/sqlalchemy_impl/project_repository.py b/packages/adapters/sqlalchemy_impl/project_repository.py index fbd4795fd..33012801b 100644 --- a/packages/adapters/sqlalchemy_impl/project_repository.py +++ b/packages/adapters/sqlalchemy_impl/project_repository.py @@ -47,9 +47,7 @@ class SQLAlchemyProjectRepository: def find_by_owner_user_id(self, owner_user_id: str) -> list[Project]: """根据所有者用户 ID 查找项目""" - models = self.session.query(ProjectModel).filter( - ProjectModel.owner_user_id == owner_user_id - ).all() + models = self.session.query(ProjectModel).filter(ProjectModel.owner_user_id == owner_user_id).all() return [self._to_entity(model) for model in models] def find_accessible_projects(self, user_id: str) -> list[Project]: @@ -57,19 +55,18 @@ class SQLAlchemyProjectRepository: from sqlalchemy import or_, cast from sqlalchemy.dialects.postgresql import JSONB - models = self.session.query(ProjectModel).filter( - or_( - ProjectModel.owner_user_id == user_id, - cast(ProjectModel.shared_users, JSONB).contains([user_id]) + models = ( + self.session.query(ProjectModel) + .filter( + or_(ProjectModel.owner_user_id == user_id, cast(ProjectModel.shared_users, JSONB).contains([user_id])) ) - ).all() + .all() + ) return [self._to_entity(model) for model in models] def count_by_owner(self, owner_user_id: str) -> int: """统计用户的项目数量""" - return self.session.query(ProjectModel).filter( - ProjectModel.owner_user_id == owner_user_id - ).count() + return self.session.query(ProjectModel).filter(ProjectModel.owner_user_id == owner_user_id).count() def delete(self, project_id: str) -> bool: """删除项目""" diff --git a/packages/adapters/sqlalchemy_impl/recipe_repository.py b/packages/adapters/sqlalchemy_impl/recipe_repository.py index 3397e4439..2687fcfc0 100644 --- a/packages/adapters/sqlalchemy_impl/recipe_repository.py +++ b/packages/adapters/sqlalchemy_impl/recipe_repository.py @@ -1,4 +1,5 @@ """SQLAlchemy implementation of RecipeRepository.""" + from __future__ import annotations from typing import List, Optional @@ -144,11 +145,7 @@ class SQLAlchemyRecipeRepository: return items def delete_items_by_recipe(self, recipe_id: str) -> int: - count = ( - self.session.query(RecipeItemModel) - .filter(RecipeItemModel.recipe_id == recipe_id) - .delete() - ) + count = self.session.query(RecipeItemModel).filter(RecipeItemModel.recipe_id == recipe_id).delete() self.session.commit() return count diff --git a/packages/adapters/sqlalchemy_impl/template_repository.py b/packages/adapters/sqlalchemy_impl/template_repository.py index a006692ff..90f4b1bfb 100644 --- a/packages/adapters/sqlalchemy_impl/template_repository.py +++ b/packages/adapters/sqlalchemy_impl/template_repository.py @@ -1,4 +1,5 @@ """SQLAlchemy implementation of TemplateRepository.""" + from __future__ import annotations from typing import List, Optional @@ -178,9 +179,7 @@ class SQLAlchemyTemplateRepository: def delete_segments_by_template(self, template_id: str) -> int: count = ( - self.session.query(TemplateSegmentModel) - .filter(TemplateSegmentModel.template_id == template_id) - .delete() + self.session.query(TemplateSegmentModel).filter(TemplateSegmentModel.template_id == template_id).delete() ) self.session.commit() return count diff --git a/packages/adapters/sqlalchemy_impl/title_library_repository.py b/packages/adapters/sqlalchemy_impl/title_library_repository.py index aeb67b944..c784c55dd 100644 --- a/packages/adapters/sqlalchemy_impl/title_library_repository.py +++ b/packages/adapters/sqlalchemy_impl/title_library_repository.py @@ -1,4 +1,5 @@ """SQLAlchemy implementation of TitleLibraryRepository.""" + from __future__ import annotations from typing import List, Optional @@ -35,10 +36,14 @@ class SQLAlchemyTitleLibraryRepository: return [self._model_to_entity(m) for m in models] def get(self, title_id: str, user_id: str) -> Optional[TitleLibraryItem]: - model = self.session.query(TitleLibraryModel).filter( - TitleLibraryModel.id == title_id, - TitleLibraryModel.user_id == user_id, - ).first() + model = ( + self.session.query(TitleLibraryModel) + .filter( + TitleLibraryModel.id == title_id, + TitleLibraryModel.user_id == user_id, + ) + .first() + ) if model is None: return None return self._model_to_entity(model) @@ -62,10 +67,14 @@ class SQLAlchemyTitleLibraryRepository: return self._model_to_entity(model) def update(self, item: TitleLibraryItem) -> TitleLibraryItem: - model = self.session.query(TitleLibraryModel).filter( - TitleLibraryModel.id == item.id, - TitleLibraryModel.user_id == item.user_id, - ).first() + model = ( + self.session.query(TitleLibraryModel) + .filter( + TitleLibraryModel.id == item.id, + TitleLibraryModel.user_id == item.user_id, + ) + .first() + ) if model is None: raise ValueError(f"TitleLibraryItem {item.id} not found") model.name = item.name @@ -80,10 +89,14 @@ class SQLAlchemyTitleLibraryRepository: return self._model_to_entity(model) def delete(self, title_id: str, user_id: str) -> bool: - model = self.session.query(TitleLibraryModel).filter( - TitleLibraryModel.id == title_id, - TitleLibraryModel.user_id == user_id, - ).first() + model = ( + self.session.query(TitleLibraryModel) + .filter( + TitleLibraryModel.id == title_id, + TitleLibraryModel.user_id == user_id, + ) + .first() + ) if model is None: return False model.is_active = False @@ -91,10 +104,14 @@ class SQLAlchemyTitleLibraryRepository: return True def count_by_user(self, user_id: str, is_active: bool = True) -> int: - return self.session.query(TitleLibraryModel).filter( - TitleLibraryModel.user_id == user_id, - TitleLibraryModel.is_active == is_active, - ).count() + return ( + self.session.query(TitleLibraryModel) + .filter( + TitleLibraryModel.user_id == user_id, + TitleLibraryModel.is_active == is_active, + ) + .count() + ) @staticmethod def _model_to_entity(model: TitleLibraryModel) -> TitleLibraryItem: diff --git a/packages/adapters/sqlalchemy_impl/voice_library_repository.py b/packages/adapters/sqlalchemy_impl/voice_library_repository.py index bfb298c95..c0875c0f3 100644 --- a/packages/adapters/sqlalchemy_impl/voice_library_repository.py +++ b/packages/adapters/sqlalchemy_impl/voice_library_repository.py @@ -1,4 +1,5 @@ """SQLAlchemy implementation of VoiceLibraryRepository.""" + from __future__ import annotations from typing import List, Optional @@ -33,10 +34,14 @@ class SQLAlchemyVoiceLibraryRepository: return [self._model_to_entity(m) for m in models] def get(self, voice_id: str, user_id: str) -> Optional[VoiceLibraryItem]: - model = self.session.query(VoiceLibraryModel).filter( - VoiceLibraryModel.id == voice_id, - VoiceLibraryModel.user_id == user_id, - ).first() + model = ( + self.session.query(VoiceLibraryModel) + .filter( + VoiceLibraryModel.id == voice_id, + VoiceLibraryModel.user_id == user_id, + ) + .first() + ) if model is None: return None return self._model_to_entity(model) @@ -64,10 +69,14 @@ class SQLAlchemyVoiceLibraryRepository: return self._model_to_entity(model) def update(self, item: VoiceLibraryItem) -> VoiceLibraryItem: - model = self.session.query(VoiceLibraryModel).filter( - VoiceLibraryModel.id == item.id, - VoiceLibraryModel.user_id == item.user_id, - ).first() + model = ( + self.session.query(VoiceLibraryModel) + .filter( + VoiceLibraryModel.id == item.id, + VoiceLibraryModel.user_id == item.user_id, + ) + .first() + ) if model is None: raise ValueError(f"VoiceLibraryItem {item.id} not found") model.name = item.name @@ -86,10 +95,14 @@ class SQLAlchemyVoiceLibraryRepository: return self._model_to_entity(model) def delete(self, voice_id: str, user_id: str) -> bool: - model = self.session.query(VoiceLibraryModel).filter( - VoiceLibraryModel.id == voice_id, - VoiceLibraryModel.user_id == user_id, - ).first() + model = ( + self.session.query(VoiceLibraryModel) + .filter( + VoiceLibraryModel.id == voice_id, + VoiceLibraryModel.user_id == user_id, + ) + .first() + ) if model is None: return False # Soft delete by setting status to deleted @@ -98,10 +111,14 @@ class SQLAlchemyVoiceLibraryRepository: return True def count_by_user(self, user_id: str) -> int: - return self.session.query(VoiceLibraryModel).filter( - VoiceLibraryModel.user_id == user_id, - VoiceLibraryModel.status != "deleted", - ).count() + return ( + self.session.query(VoiceLibraryModel) + .filter( + VoiceLibraryModel.user_id == user_id, + VoiceLibraryModel.status != "deleted", + ) + .count() + ) @staticmethod def _model_to_entity(model: VoiceLibraryModel) -> VoiceLibraryItem: diff --git a/packages/adapters/sqlite_tracker/__init__.py b/packages/adapters/sqlite_tracker/__init__.py index 8ef8ef9d3..ef2cad630 100644 --- a/packages/adapters/sqlite_tracker/__init__.py +++ b/packages/adapters/sqlite_tracker/__init__.py @@ -1,6 +1,5 @@ """SQLite Tracker Adapter""" - __all__ = [ "SQLiteTaskRepository", "SQLiteMilestoneRepository", diff --git a/packages/application/auth/jwt_handler.py b/packages/application/auth/jwt_handler.py index 68d185949..367222fe3 100644 --- a/packages/application/auth/jwt_handler.py +++ b/packages/application/auth/jwt_handler.py @@ -6,7 +6,7 @@ JWT 处理器委托层 使用方式: from packages.application.auth.jwt_handler import JWTHandler, get_jwt_handler - + jwt_handler = JWTHandler(secret_key="") token = jwt_handler.create_access_token(user_id="user123", role="admin") payload = jwt_handler.verify_access_token(token) @@ -21,7 +21,7 @@ from packages.application.auth.jwt_service import JWTConfig, JWTService, TokenTy class JWTHandler: """ JWT 处理器委托类 - + 委托给 packages.domain.auth.jwt_service.JWTService 进行实际的 JWT 操作, 此层仅负责配置和封装,不直接依赖 jwt 库。 """ @@ -29,7 +29,7 @@ class JWTHandler: def __init__(self, secret_key: str, algorithm: str = "HS256", access_token_expire_minutes: int = 30): """ 初始化 JWT 处理器 - + Args: secret_key: JWT 签名密钥(必须从环境变量或配置注入) algorithm: 加密算法,默认 HS256 @@ -68,13 +68,13 @@ class JWTHandler: def verify_access_token(self, token: str) -> Dict[str, Any]: """ 验证 access_token - + Args: token: JWT Token 字符串 - + Returns: Token payload - + Raises: ExpiredSignatureError: Token 已过期 ValueError: Token 类型不是 access @@ -84,10 +84,10 @@ class JWTHandler: def verify_token(self, token: str) -> Dict[str, Any]: """ 验证任意 Token - + Args: token: JWT Token 字符串 - + Returns: Token payload """ @@ -105,12 +105,12 @@ def configure_jwt_handler( ) -> JWTHandler: """ 配置全局 JWT 处理器 - + Args: secret_key: JWT 签名密钥 algorithm: 加密算法 access_token_expire_minutes: Access Token 过期时间(分钟) - + Returns: 配置好的 JWTHandler 实例 """ @@ -126,15 +126,13 @@ def configure_jwt_handler( def get_jwt_handler() -> JWTHandler: """ 获取全局 JWT 处理器 - + Returns: JWTHandler 实例 - + Raises: RuntimeError: 如果尚未配置 JWT 处理器 """ if _default_handler is None: - raise RuntimeError( - "JWT handler not configured. Call configure_jwt_handler() first." - ) + raise RuntimeError("JWT handler not configured. Call configure_jwt_handler() first.") return _default_handler diff --git a/packages/application/auth/jwt_service.py b/packages/application/auth/jwt_service.py index 8b1db5930..f3c87539b 100644 --- a/packages/application/auth/jwt_service.py +++ b/packages/application/auth/jwt_service.py @@ -30,9 +30,7 @@ class JWTConfig: ValueError: 如果 secret_key 为空或包含不安全默认值 """ if not secret_key or secret_key.strip() == "": - raise ValueError( # noqa: E501 - "JWT secret_key must be provided and cannot be empty" - ) + raise ValueError("JWT secret_key must be provided and cannot be empty") # noqa: E501 insecure_defaults = [ "your-secret-key-change-in-production", @@ -43,8 +41,7 @@ class JWTConfig: ] if secret_key.lower() in [d.lower() for d in insecure_defaults]: raise ValueError( # noqa: E501 - f"JWT secret_key '{secret_key}' is insecure. " - "Please provide a strong random secret." + f"JWT secret_key '{secret_key}' is insecure. " "Please provide a strong random secret." ) self.SECRET_KEY: str = secret_key @@ -90,9 +87,7 @@ class JWTService: JWT Token 字符串 """ now = datetime.utcnow() - expire = now + timedelta( # noqa: E501 - minutes=self.config.ACCESS_TOKEN_EXPIRE_MINUTES - ) + expire = now + timedelta(minutes=self.config.ACCESS_TOKEN_EXPIRE_MINUTES) # noqa: E501 payload = { "sub": user_id, # subject (用户ID) @@ -105,9 +100,7 @@ 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: """ @@ -131,9 +124,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]: """ @@ -223,10 +214,12 @@ class JWTService: # Lazy singleton - created with settings on first access _jwt_service_instance = None + def _get_jwt_service(): global _jwt_service_instance if _jwt_service_instance is None: from app.config import settings + kw = dict(secret_key=settings.JWT_SECRET_KEY) if hasattr(settings, "JWT_ALGORITHM"): kw["algorithm"] = settings.JWT_ALGORITHM @@ -237,9 +230,10 @@ def _get_jwt_service(): _jwt_service_instance = JWTService(JWTConfig(**kw)) return _jwt_service_instance + class _JWTServiceProxy: def __getattr__(self, name): return getattr(_get_jwt_service(), name) -jwt_service = _JWTServiceProxy() +jwt_service = _JWTServiceProxy() diff --git a/packages/application/auth/login_use_case.py b/packages/application/auth/login_use_case.py index fec87e793..1f26acc80 100755 --- a/packages/application/auth/login_use_case.py +++ b/packages/application/auth/login_use_case.py @@ -247,13 +247,13 @@ class RefreshTokenUseCase: def _find_session_by_refresh_token(self, refresh_token: str) -> Optional[dict]: """ 通过 refresh_token 查找 session - + 使用 Redis 中的反向索引 (refresh_token -> session_id) 快速查找 session。 反向索引在 save_session 时创建,确保了 O(1) 的查找复杂度。 - + Args: refresh_token: 刷新令牌 - + Returns: Session 数据字典,包含 session_id, user_id 等信息;如果不存在返回 None """ diff --git a/packages/application/auth/password_handler.py b/packages/application/auth/password_handler.py index aeae77b98..4d97704bc 100644 --- a/packages/application/auth/password_handler.py +++ b/packages/application/auth/password_handler.py @@ -6,7 +6,7 @@ 使用方式: from packages.application.auth.password_handler import PasswordHandler, get_password_handler - + password_handler = PasswordHandler() hashed = password_handler.hash_password("my_secure_password") is_valid = password_handler.verify_password("my_secure_password", hashed) @@ -20,7 +20,7 @@ from packages.application.auth.password_hasher import PasswordHasher, PasswordVa class PasswordHandler: """ 密码处理器委托类 - + 委托给 packages.domain.auth.password_hasher 进行实际的密码哈希操作, 此层仅负责配置和封装,不直接依赖 bcrypt 库。 """ @@ -28,7 +28,7 @@ class PasswordHandler: def __init__(self, rounds: int = 12): """ 初始化密码处理器 - + Args: rounds: bcrypt cost factor(默认 12,推荐范围 10-14) """ @@ -44,13 +44,13 @@ class PasswordHandler: def hash_password(self, password: str) -> str: """ 哈希密码 - + Args: password: 明文密码 - + Returns: bcrypt 哈希字符串 - + Raises: ValueError: 密码为空 """ @@ -59,11 +59,11 @@ class PasswordHandler: def verify_password(self, password: str, hashed_password: str) -> bool: """ 验证密码 - + Args: password: 明文密码 hashed_password: 存储的哈希密码 - + Returns: True 如果密码正确,否则 False """ @@ -72,10 +72,10 @@ class PasswordHandler: def needs_rehash(self, hashed_password: str) -> bool: """ 检查哈希是否需要重新计算 - + Args: hashed_password: 存储的哈希密码 - + Returns: True 如果需要重新哈希 """ @@ -84,10 +84,10 @@ class PasswordHandler: def validate_strength(self, password: str) -> Tuple[bool, Optional[str]]: """ 验证密码强度 - + Args: password: 明文密码 - + Returns: (是否有效, 错误信息) """ @@ -101,10 +101,10 @@ _default_handler: Optional[PasswordHandler] = None def configure_password_handler(rounds: int = 12) -> PasswordHandler: """ 配置全局密码处理器 - + Args: rounds: bcrypt cost factor - + Returns: 配置好的 PasswordHandler 实例 """ @@ -116,7 +116,7 @@ def configure_password_handler(rounds: int = 12) -> PasswordHandler: def get_password_handler() -> PasswordHandler: """ 获取全局密码处理器 - + Returns: PasswordHandler 实例 """ diff --git a/packages/application/duplication.py b/packages/application/duplication.py index 4879a509a..66a5950e6 100644 --- a/packages/application/duplication.py +++ b/packages/application/duplication.py @@ -1,4 +1,5 @@ """查重应用层用例。""" + from __future__ import annotations import logging diff --git a/packages/application/recipe/commands.py b/packages/application/recipe/commands.py index fe7535fd1..97757a2d4 100644 --- a/packages/application/recipe/commands.py +++ b/packages/application/recipe/commands.py @@ -1,4 +1,5 @@ """Recipe commands.""" + from __future__ import annotations from dataclasses import dataclass, field diff --git a/packages/application/recipe/use_cases.py b/packages/application/recipe/use_cases.py index 0eba4bfe3..bb5746217 100644 --- a/packages/application/recipe/use_cases.py +++ b/packages/application/recipe/use_cases.py @@ -1,4 +1,5 @@ """Recipe use cases.""" + from __future__ import annotations import uuid @@ -26,6 +27,7 @@ class FeatureDisabledError(Exception): @dataclass class MissingAssetWarning: """使用配方时缺失的素材警告""" + item_type: str item_id: str position: int @@ -144,6 +146,7 @@ class DeleteRecipeUseCase: @dataclass class UseRecipeResult: """使用配方的结果""" + recipe: Recipe warnings: List[MissingAssetWarning] @@ -166,9 +169,7 @@ class UseRecipeUseCase: FeatureScope.RECIPE_REUSE, user_plan=user_plan, ): - raise FeatureDisabledError( - "配方复用功能仅对基础版和高级版用户开放" - ) + raise FeatureDisabledError("配方复用功能仅对基础版和高级版用户开放") # 2. 获取配方 recipe = self.repository.get(recipe_id, user_id) diff --git a/packages/application/template/commands.py b/packages/application/template/commands.py index 0fdc71e8e..7fe4901ca 100644 --- a/packages/application/template/commands.py +++ b/packages/application/template/commands.py @@ -1,4 +1,5 @@ """Template commands.""" + from __future__ import annotations from dataclasses import dataclass, field diff --git a/packages/application/template/use_cases.py b/packages/application/template/use_cases.py index f43cb0122..1f69c2f11 100644 --- a/packages/application/template/use_cases.py +++ b/packages/application/template/use_cases.py @@ -1,4 +1,5 @@ """Template use cases.""" + from __future__ import annotations import uuid @@ -22,6 +23,7 @@ class NotFoundError(Exception): class ValidationError(Exception): """业务规则校验失败.""" + pass @@ -32,6 +34,7 @@ VALID_MATERIAL_TYPES = {"人物", "场景"} @dataclass class GenerateWarning: """生成时的警告信息.""" + code: str # voiceover_duration_mismatch / missing_material_type / ... message: str details: dict = field(default_factory=dict) @@ -40,6 +43,7 @@ class GenerateWarning: @dataclass class ValidateResult: """模板校验结果.""" + template: Template warnings: List[GenerateWarning] = field(default_factory=list) @@ -190,9 +194,7 @@ class ValidateTemplateUseCase: # 业务规则 1: one_take 必须恰好 1 个片段 if template.mode == EditingMode.ONE_TAKE.value: if len(template.segments) != 1: - raise ValidationError( - f"一镜到底模式必须恰好有 1 个片段,当前有 {len(template.segments)} 个" - ) + raise ValidationError(f"一镜到底模式必须恰好有 1 个片段,当前有 {len(template.segments)} 个") # 业务规则 2: voice_over 每个片段必须有 material_type if template.mode == EditingMode.VOICE_OVER.value: @@ -207,19 +209,21 @@ class ValidateTemplateUseCase: if command.voiceover_duration is not None and template.estimated_duration > 0: ratio = command.voiceover_duration / template.estimated_duration if ratio < 0.7 or ratio > 1.3: - warnings.append(GenerateWarning( - code="voiceover_duration_mismatch", - message=( - f"配音时长 ({command.voiceover_duration:.1f}s) " - f"与预估时长 ({template.estimated_duration:.1f}s) " - f"偏差超过 ±30%,可能影响剪辑效果" - ), - details={ - "voiceover_duration": command.voiceover_duration, - "estimated_duration": template.estimated_duration, - "ratio": round(ratio, 3), - }, - )) + warnings.append( + GenerateWarning( + code="voiceover_duration_mismatch", + message=( + f"配音时长 ({command.voiceover_duration:.1f}s) " + f"与预估时长 ({template.estimated_duration:.1f}s) " + f"偏差超过 ±30%,可能影响剪辑效果" + ), + details={ + "voiceover_duration": command.voiceover_duration, + "estimated_duration": template.estimated_duration, + "ratio": round(ratio, 3), + }, + ) + ) return ValidateResult(template=template, warnings=warnings) diff --git a/packages/application/title_library/__init__.py b/packages/application/title_library/__init__.py index 384a9411e..df83c0eff 100644 --- a/packages/application/title_library/__init__.py +++ b/packages/application/title_library/__init__.py @@ -1,4 +1,5 @@ """Title library application module.""" + from packages.application.title_library.use_cases import ( CreateTitleLibraryUseCase, DeleteTitleLibraryUseCase, diff --git a/packages/application/title_library/commands.py b/packages/application/title_library/commands.py index d65acbf2e..0f4cd7012 100644 --- a/packages/application/title_library/commands.py +++ b/packages/application/title_library/commands.py @@ -1,4 +1,5 @@ """Title library commands.""" + from __future__ import annotations from dataclasses import dataclass, field diff --git a/packages/application/title_library/use_cases.py b/packages/application/title_library/use_cases.py index f6fe8d4e4..4306b3402 100644 --- a/packages/application/title_library/use_cases.py +++ b/packages/application/title_library/use_cases.py @@ -1,4 +1,5 @@ """Title library use cases.""" + from __future__ import annotations import uuid diff --git a/packages/application/voice_library/__init__.py b/packages/application/voice_library/__init__.py index ef25f3343..521682be3 100644 --- a/packages/application/voice_library/__init__.py +++ b/packages/application/voice_library/__init__.py @@ -1,4 +1,5 @@ """Voice library application module.""" + from packages.application.voice_library.use_cases import ( CreateVoiceLibraryUseCase, DeleteVoiceLibraryUseCase, diff --git a/packages/application/voice_library/commands.py b/packages/application/voice_library/commands.py index 210a99f54..1cfde76e6 100644 --- a/packages/application/voice_library/commands.py +++ b/packages/application/voice_library/commands.py @@ -1,4 +1,5 @@ """Voice library commands.""" + from __future__ import annotations from dataclasses import dataclass, field diff --git a/packages/application/voice_library/use_cases.py b/packages/application/voice_library/use_cases.py index 2c3ae5b58..b02a0c76b 100644 --- a/packages/application/voice_library/use_cases.py +++ b/packages/application/voice_library/use_cases.py @@ -1,4 +1,5 @@ """Voice library use cases.""" + from __future__ import annotations import uuid diff --git a/packages/domain/duplication.py b/packages/domain/duplication.py index d83785e13..cd5e8dcf8 100644 --- a/packages/domain/duplication.py +++ b/packages/domain/duplication.py @@ -1,4 +1,5 @@ """查重记录领域实体。""" + from __future__ import annotations from dataclasses import dataclass, field diff --git a/packages/domain/editing_mode.py b/packages/domain/editing_mode.py index 9acc85a81..0932d28a3 100644 --- a/packages/domain/editing_mode.py +++ b/packages/domain/editing_mode.py @@ -5,7 +5,8 @@ from enum import StrEnum class EditingMode(StrEnum): """剪辑模式枚举""" - ONE_TAKE = "one_take" # 顺序拼接模式 - PIP = "pip" # 画中画模式 - VOICE_OVER = "voice_over" # 口播+B-roll模式 - VOICE_PIP = "voice_pip" # 口播+画中画组合模式 + + ONE_TAKE = "one_take" # 顺序拼接模式 + PIP = "pip" # 画中画模式 + VOICE_OVER = "voice_over" # 口播+B-roll模式 + VOICE_PIP = "voice_pip" # 口播+画中画组合模式 diff --git a/packages/domain/entities.py b/packages/domain/entities.py index 7735c0c96..88b33178c 100644 --- a/packages/domain/entities.py +++ b/packages/domain/entities.py @@ -248,10 +248,3 @@ class IngestJob: library_id=library_id.strip(), storage_key=storage_key.strip(), ) - - - - - - - diff --git a/packages/domain/quota.py b/packages/domain/quota.py index 8397f469b..3f2f79e9c 100644 --- a/packages/domain/quota.py +++ b/packages/domain/quota.py @@ -18,23 +18,25 @@ from typing import Dict, List, Optional class QuotaDimension(str, Enum): """配额维度 - 所有可量化的资源限制""" - STORAGE_GB = "storage_gb" # 存储空间 (GB) - VIDEOS_PER_MONTH = "videos_per_month" # 每月生成视频数 - MAX_CONCURRENT = "max_concurrent" # 最大并发任务数 - MAX_TEMPLATES = "max_templates" # 最大模板数 - MAX_TITLES = "max_titles" # 最大标题库条目数 - MAX_VOICEOVERS = "max_voiceovers" # 最大配音库条目数 - AI_VOICE_ENABLED = "ai_voice_enabled" # AI 配音是否可用 (0/1) + + STORAGE_GB = "storage_gb" # 存储空间 (GB) + VIDEOS_PER_MONTH = "videos_per_month" # 每月生成视频数 + MAX_CONCURRENT = "max_concurrent" # 最大并发任务数 + MAX_TEMPLATES = "max_templates" # 最大模板数 + MAX_TITLES = "max_titles" # 最大标题库条目数 + MAX_VOICEOVERS = "max_voiceovers" # 最大配音库条目数 + AI_VOICE_ENABLED = "ai_voice_enabled" # AI 配音是否可用 (0/1) # 以下维度由扩展模块注册,初始配额为 0(由模块注册时填充) - AI_VOICE_CREDITS = "ai_voice_credits" # AI 配音积分(每月) - BATCH_EXPORT_ENABLED = "batch_export_enabled" # 批量导出 - MULTI_PLATFORM_ENABLED = "multi_platform_enabled" # 多平台发布 - DEDUP_REPORT_ENABLED = "dedup_report_enabled" # 去重检测报告 + AI_VOICE_CREDITS = "ai_voice_credits" # AI 配音积分(每月) + BATCH_EXPORT_ENABLED = "batch_export_enabled" # 批量导出 + MULTI_PLATFORM_ENABLED = "multi_platform_enabled" # 多平台发布 + DEDUP_REPORT_ENABLED = "dedup_report_enabled" # 去重检测报告 @dataclass class QuotaTier: """一个套餐等级的配额定义""" + name: str limits: Dict[str, float] = field(default_factory=dict) @@ -102,15 +104,17 @@ QUOTA_TIERS: Dict[str, QuotaTier] = { class QuotaWarningLevel: """配额告警级别""" - NORMAL = "normal" # 使用量 < 80% - WARNING = "warning" # 80% <= 使用量 < 100% - CRITICAL = "critical" # 95% <= 使用量 < 100% - EXCEEDED = "exceeded" # 使用量 >= 100% + + NORMAL = "normal" # 使用量 < 80% + WARNING = "warning" # 80% <= 使用量 < 100% + CRITICAL = "critical" # 95% <= 使用量 < 100% + EXCEEDED = "exceeded" # 使用量 >= 100% @dataclass class QuotaCheckResult: """配额检查结果""" + allowed: bool dimension: str limit: float @@ -236,10 +240,7 @@ class QuotaChecker: usage: Dict[str, float], ) -> List[QuotaCheckResult]: """批量检查多个维度的配额""" - return [ - self.check(plan_name, dim, used) - for dim, used in usage.items() - ] + return [self.check(plan_name, dim, used) for dim, used in usage.items()] @staticmethod def _compute_warning_level(used: float, limit: float) -> str: diff --git a/packages/domain/recipe.py b/packages/domain/recipe.py index c43c550ca..1b11134d7 100644 --- a/packages/domain/recipe.py +++ b/packages/domain/recipe.py @@ -1,4 +1,5 @@ """Recipe domain entities.""" + from __future__ import annotations from dataclasses import dataclass, field @@ -9,6 +10,7 @@ from typing import List @dataclass class RecipeItem: """配方中的单个素材/标题/配音项""" + id: str recipe_id: str item_type: str # asset / title / voice @@ -20,6 +22,7 @@ class RecipeItem: @dataclass class Recipe: """配方 — 一次「一键生成」的完整参数组合""" + id: str user_id: str name: str diff --git a/packages/domain/template.py b/packages/domain/template.py index c9e27bbc6..85d17d829 100644 --- a/packages/domain/template.py +++ b/packages/domain/template.py @@ -1,4 +1,5 @@ """Template domain entities — 剪辑计划模板.""" + from __future__ import annotations from dataclasses import dataclass, field @@ -9,6 +10,7 @@ from typing import List, Optional @dataclass class TemplateSegment: """模板中的单个片段.""" + id: str template_id: str segment_order: int @@ -22,6 +24,7 @@ class TemplateSegment: @dataclass class Template: """剪辑计划模板.""" + id: str user_id: str name: str @@ -41,6 +44,7 @@ class Template: @dataclass class TemplateCategory: """模板分类.""" + id: str user_id: str name: str diff --git a/packages/domain/title_library.py b/packages/domain/title_library.py index 65d75c5f6..f1e5ad3db 100644 --- a/packages/domain/title_library.py +++ b/packages/domain/title_library.py @@ -1,4 +1,5 @@ """Title library domain entity.""" + from __future__ import annotations from dataclasses import dataclass, field @@ -9,6 +10,7 @@ from typing import List @dataclass class TitleLibraryItem: """标题库条目""" + id: str user_id: str name: str diff --git a/packages/domain/voice_library.py b/packages/domain/voice_library.py index 683c38558..2f8d10c2f 100644 --- a/packages/domain/voice_library.py +++ b/packages/domain/voice_library.py @@ -1,4 +1,5 @@ """Voice library domain entity.""" + from __future__ import annotations from dataclasses import dataclass, field @@ -9,6 +10,7 @@ from typing import List, Optional @dataclass class VoiceLibraryItem: """配音库条目""" + id: str user_id: str name: str diff --git a/packages/infrastructure/feature_flags.py b/packages/infrastructure/feature_flags.py index 84e0cf681..076a62109 100644 --- a/packages/infrastructure/feature_flags.py +++ b/packages/infrastructure/feature_flags.py @@ -32,6 +32,7 @@ logger = logging.getLogger(__name__) @dataclass class FeatureFlag: """单个 Feature Flag 的定义""" + name: str description: str = "" # 全局开关,默认 True(启用) @@ -64,6 +65,7 @@ class FeatureFlag: class FeatureScope: """Feature Flag 名称常量,避免硬编码字符串""" + AI_VOICE_GENERATION = "ai_voice_generation" DEDUPLICATION_REPORT = "deduplication_report" BATCH_EXPORT = "batch_export" @@ -179,10 +181,7 @@ class FeatureFlags: def get_enabled_for_plan(self, plan: str) -> list[str]: """获取指定套餐下所有启用的功能名称""" - return [ - name for name, flag in self._flags.items() - if flag.is_enabled(user_plan=plan) - ] + return [name for name, flag in self._flags.items() if flag.is_enabled(user_plan=plan)] # 全局单例 diff --git a/packages/infrastructure/module_registry.py b/packages/infrastructure/module_registry.py index 5ce1843f5..9495a8e50 100644 --- a/packages/infrastructure/module_registry.py +++ b/packages/infrastructure/module_registry.py @@ -21,10 +21,11 @@ logger = logging.getLogger(__name__) class ModuleStatus(str, Enum): """模块运行状态""" - REGISTERED = "registered" # 已注册,未激活 - ACTIVE = "active" # 已激活,可用 - DISABLED = "disabled" # 已禁用(管理员/Feature Flag 控制) - ERROR = "error" # 注册或初始化出错 + + REGISTERED = "registered" # 已注册,未激活 + ACTIVE = "active" # 已激活,可用 + DISABLED = "disabled" # 已禁用(管理员/Feature Flag 控制) + ERROR = "error" # 注册或初始化出错 @dataclass @@ -34,9 +35,10 @@ class QuotaRule: 描述该模块消耗哪些配额维度,以及每个操作消耗多少。 例如:AI 配音模块每生成一条配音消耗 1 个 ai_voice_credit。 """ - dimension: str # 配额维度名,如 "ai_voice_credits", "storage_gb" - per_operation: float # 每次操作消耗量 - description: str = "" # 人类可读描述 + + dimension: str # 配额维度名,如 "ai_voice_credits", "storage_gb" + per_operation: float # 每次操作消耗量 + description: str = "" # 人类可读描述 @dataclass @@ -45,10 +47,11 @@ class ModuleCapability: 能力是模块对外暴露的可调用功能单元。 """ - name: str # 能力名,如 "generate_voice" - description: str = "" # 人类可读描述 + + name: str # 能力名,如 "generate_voice" + description: str = "" # 人类可读描述 quota_rules: List[QuotaRule] = field(default_factory=list) # 该能力消耗的配额规则 - metadata: Dict[str, Any] = field(default_factory=dict) # 扩展元数据 + metadata: Dict[str, Any] = field(default_factory=dict) # 扩展元数据 @dataclass @@ -57,13 +60,14 @@ class Module: 每个扩展模块通过 Module 描述自身,注册到 ModuleRegistry。 """ - name: str # 模块唯一标识,如 "ai_voice" - version: str = "1.0.0" # 模块版本 - description: str = "" # 人类可读描述 + + name: str # 模块唯一标识,如 "ai_voice" + version: str = "1.0.0" # 模块版本 + description: str = "" # 人类可读描述 capabilities: List[ModuleCapability] = field(default_factory=list) - dependencies: List[str] = field(default_factory=list) # 依赖的其他模块名 + dependencies: List[str] = field(default_factory=list) # 依赖的其他模块名 status: ModuleStatus = ModuleStatus.REGISTERED - config: Dict[str, Any] = field(default_factory=dict) # 模块配置 + config: Dict[str, Any] = field(default_factory=dict) # 模块配置 _init_func: Optional[Callable] = field(default=None, repr=False) # 初始化回调 def activate(self) -> None: @@ -144,14 +148,9 @@ class ModuleRegistry: raise KeyError(f"Module '{name}' not found") # 检查是否有其他模块依赖它 - dependents = [ - m.name for m in self._modules.values() - if name in m.dependencies and m.name != name - ] + dependents = [m.name for m in self._modules.values() if name in m.dependencies and m.name != name] if dependents: - raise ValueError( - f"Cannot unregister module '{name}': depended on by {dependents}" - ) + raise ValueError(f"Cannot unregister module '{name}': depended on by {dependents}") del self._modules[name] logger.info(f"Module '{name}' unregistered") diff --git a/packages/ports/duplication_repository.py b/packages/ports/duplication_repository.py index 40e9bfcb9..9cf8aac0e 100644 --- a/packages/ports/duplication_repository.py +++ b/packages/ports/duplication_repository.py @@ -1,4 +1,5 @@ """查重记录仓库端口(Protocol)。""" + from __future__ import annotations from typing import Protocol diff --git a/packages/ports/recipe_repository.py b/packages/ports/recipe_repository.py index a17f7cff2..fb6ceb347 100644 --- a/packages/ports/recipe_repository.py +++ b/packages/ports/recipe_repository.py @@ -1,4 +1,5 @@ """Recipe repository port.""" + from __future__ import annotations from typing import List, Optional, Protocol @@ -15,29 +16,20 @@ class RecipeRepository(Protocol): *, skip: int = 0, limit: int = 50, - ) -> List[Recipe]: - ... + ) -> List[Recipe]: ... - def get(self, recipe_id: str, user_id: str) -> Optional[Recipe]: - ... + def get(self, recipe_id: str, user_id: str) -> Optional[Recipe]: ... - def create(self, recipe: Recipe) -> Recipe: - ... + def create(self, recipe: Recipe) -> Recipe: ... - def update(self, recipe: Recipe) -> Recipe: - ... + def update(self, recipe: Recipe) -> Recipe: ... - def delete(self, recipe_id: str, user_id: str) -> bool: - ... + def delete(self, recipe_id: str, user_id: str) -> bool: ... - def count_by_user(self, user_id: str, is_active: bool = True) -> int: - ... + def count_by_user(self, user_id: str, is_active: bool = True) -> int: ... - def list_items(self, recipe_id: str) -> List[RecipeItem]: - ... + def list_items(self, recipe_id: str) -> List[RecipeItem]: ... - def create_items(self, items: List[RecipeItem]) -> List[RecipeItem]: - ... + def create_items(self, items: List[RecipeItem]) -> List[RecipeItem]: ... - def delete_items_by_recipe(self, recipe_id: str) -> int: - ... + def delete_items_by_recipe(self, recipe_id: str) -> int: ... diff --git a/packages/ports/template_repository.py b/packages/ports/template_repository.py index a93a7d4e7..6b59372be 100644 --- a/packages/ports/template_repository.py +++ b/packages/ports/template_repository.py @@ -1,4 +1,5 @@ """Template repository port (Protocol).""" + from __future__ import annotations from typing import List, Optional, Protocol diff --git a/packages/ports/title_library_repository.py b/packages/ports/title_library_repository.py index e74a50efc..22b020544 100644 --- a/packages/ports/title_library_repository.py +++ b/packages/ports/title_library_repository.py @@ -1,4 +1,5 @@ """Title library repository port.""" + from __future__ import annotations from typing import List, Optional, Protocol @@ -17,20 +18,14 @@ class TitleLibraryRepository(Protocol): is_active: bool = True, skip: int = 0, limit: int = 50, - ) -> List[TitleLibraryItem]: - ... + ) -> List[TitleLibraryItem]: ... - def get(self, title_id: str, user_id: str) -> Optional[TitleLibraryItem]: - ... + def get(self, title_id: str, user_id: str) -> Optional[TitleLibraryItem]: ... - def create(self, item: TitleLibraryItem) -> TitleLibraryItem: - ... + def create(self, item: TitleLibraryItem) -> TitleLibraryItem: ... - def update(self, item: TitleLibraryItem) -> TitleLibraryItem: - ... + def update(self, item: TitleLibraryItem) -> TitleLibraryItem: ... - def delete(self, title_id: str, user_id: str) -> bool: - ... + def delete(self, title_id: str, user_id: str) -> bool: ... - def count_by_user(self, user_id: str, is_active: bool = True) -> int: - ... + def count_by_user(self, user_id: str, is_active: bool = True) -> int: ... diff --git a/packages/ports/voice_library_repository.py b/packages/ports/voice_library_repository.py index f4922cb1a..2e97d3730 100644 --- a/packages/ports/voice_library_repository.py +++ b/packages/ports/voice_library_repository.py @@ -1,4 +1,5 @@ """Voice library repository port.""" + from __future__ import annotations from typing import List, Optional, Protocol @@ -16,20 +17,14 @@ class VoiceLibraryRepository(Protocol): status: Optional[str] = None, skip: int = 0, limit: int = 50, - ) -> List[VoiceLibraryItem]: - ... + ) -> List[VoiceLibraryItem]: ... - def get(self, voice_id: str, user_id: str) -> Optional[VoiceLibraryItem]: - ... + def get(self, voice_id: str, user_id: str) -> Optional[VoiceLibraryItem]: ... - def create(self, item: VoiceLibraryItem) -> VoiceLibraryItem: - ... + def create(self, item: VoiceLibraryItem) -> VoiceLibraryItem: ... - def update(self, item: VoiceLibraryItem) -> VoiceLibraryItem: - ... + def update(self, item: VoiceLibraryItem) -> VoiceLibraryItem: ... - def delete(self, voice_id: str, user_id: str) -> bool: - ... + def delete(self, voice_id: str, user_id: str) -> bool: ... - def count_by_user(self, user_id: str) -> int: - ... + def count_by_user(self, user_id: str) -> int: ... diff --git a/scripts/smoke_test.py b/scripts/smoke_test.py index 5d8675709..d0c052f02 100644 --- a/scripts/smoke_test.py +++ b/scripts/smoke_test.py @@ -4,6 +4,7 @@ 用法: python3 smoke_test.py [--email EMAIL] [--password PASSWORD] [--json] 示例: python3 smoke_test.py https://saas-api.xiaoxiajianji.com --email test@example.com --password test123 --json """ + import argparse import json import sys @@ -13,19 +14,29 @@ import urllib.error import ssl CORE_ENDPOINTS = [ - {"name": "upload/direct/prepare", "method": "POST", "path": "/api/v1/upload/direct/prepare", - "body": {"project_id": "smoke-test", "file_name": "test.mp4", "file_size": 1024, "content_type": "video/mp4"}, - "expect": [200, 401, 422]}, - {"name": "upload/chunk/init", "method": "POST", "path": "/api/v1/upload/chunk/init", - "body": {"project_id": "smoke-test", "file_name": "test.mp4", "file_size": 1024000, "total_chunks": 2}, - "expect": [200, 401, 422]}, - {"name": "dashboard/overview", "method": "GET", "path": "/api/v1/dashboard/overview", - "expect": [200, 401]}, - {"name": "assets", "method": "GET", "path": "/api/v1/assets?library_id=smoke-test", - "expect": [200, 401]}, - {"name": "generation/tasks", "method": "POST", "path": "/api/v1/generation/tasks", - "body": {}, - "expect": [200, 401, 422]}, + { + "name": "upload/direct/prepare", + "method": "POST", + "path": "/api/v1/upload/direct/prepare", + "body": {"project_id": "smoke-test", "file_name": "test.mp4", "file_size": 1024, "content_type": "video/mp4"}, + "expect": [200, 401, 422], + }, + { + "name": "upload/chunk/init", + "method": "POST", + "path": "/api/v1/upload/chunk/init", + "body": {"project_id": "smoke-test", "file_name": "test.mp4", "file_size": 1024000, "total_chunks": 2}, + "expect": [200, 401, 422], + }, + {"name": "dashboard/overview", "method": "GET", "path": "/api/v1/dashboard/overview", "expect": [200, 401]}, + {"name": "assets", "method": "GET", "path": "/api/v1/assets?library_id=smoke-test", "expect": [200, 401]}, + { + "name": "generation/tasks", + "method": "POST", + "path": "/api/v1/generation/tasks", + "body": {}, + "expect": [200, 401, 422], + }, ] @@ -34,14 +45,14 @@ def make_request(base_url, endpoint, token=None): headers = {"Content-Type": "application/json"} if token: headers["Authorization"] = f"Bearer {token}" - + data = json.dumps(endpoint.get("body", {})).encode() if endpoint.get("body") is not None else None req = urllib.request.Request(url, data=data, headers=headers, method=endpoint["method"]) - + ctx = ssl.create_default_context() ctx.check_hostname = False ctx.verify_mode = ssl.CERT_NONE - + try: start = time.time() resp = urllib.request.urlopen(req, timeout=15, context=ctx) @@ -78,15 +89,15 @@ def login(base_url, email, password): def run_smoke_test(base_url, email=None, password=None, output_json=False): base_url = base_url.rstrip("/") token = None - + if email and password: token = login(base_url, email, password) if not output_json: print(f"{'✅ 登录成功' if token else '⚠️ 登录失败,将以未认证模式测试'}") - + results = [] all_passed = True - + for ep in CORE_ENDPOINTS: result = make_request(base_url, ep, token) passed = result["status"] in ep["expect"] and result["error"] is None @@ -94,24 +105,26 @@ def run_smoke_test(base_url, email=None, password=None, output_json=False): if is_5xx: passed = False all_passed = False - - results.append({ - "name": ep["name"], - "path": ep["path"], - "status": result["status"], - "elapsed_ms": result["elapsed_ms"], - "passed": passed, - "error": result["error"], - "is_5xx": is_5xx - }) - + + results.append( + { + "name": ep["name"], + "path": ep["path"], + "status": result["status"], + "elapsed_ms": result["elapsed_ms"], + "passed": passed, + "error": result["error"], + "is_5xx": is_5xx, + } + ) + if not output_json: icon = "✅" if passed else "❌" print(f" {icon} {ep['name']}: {result['status']} ({result['elapsed_ms']}ms)") - + if output_json: print(json.dumps({"success": all_passed, "results": results, "base_url": base_url}, indent=2)) - + return 0 if all_passed else 1 @@ -122,17 +135,17 @@ def main(): parser.add_argument("--password", help="登录密码") parser.add_argument("--json", action="store_true", help="JSON 格式输出") args = parser.parse_args() - + if not args.json: print(f"\n🔍 冒烟测试: {args.base_url}") print("-" * 50) - + exit_code = run_smoke_test(args.base_url, args.email, args.password, args.json) - + if not args.json: print("-" * 50) print(f"{'✅ 全部通过' if exit_code == 0 else '❌ 存在失败端点'}\n") - + sys.exit(exit_code) diff --git a/tests/integration/test_auth.py b/tests/integration/test_auth.py index f72c8c018..f9fc4e4f5 100644 --- a/tests/integration/test_auth.py +++ b/tests/integration/test_auth.py @@ -9,7 +9,6 @@ from fastapi.testclient import TestClient from apps.api.main import app - client = TestClient(app) diff --git a/tests/integration/test_duplication_upload_error_handling.py b/tests/integration/test_duplication_upload_error_handling.py index 87d60a8bf..ace0cc237 100644 --- a/tests/integration/test_duplication_upload_error_handling.py +++ b/tests/integration/test_duplication_upload_error_handling.py @@ -8,6 +8,7 @@ 覆盖端点:POST /upload(查重上传) """ + from __future__ import annotations import io @@ -22,11 +23,11 @@ import pytest from fastapi import FastAPI from fastapi.testclient import TestClient - # --------------------------------------------------------------------------- # 1. Mock 项目内部模块 # --------------------------------------------------------------------------- + def _install_mocks(): """安装所有必需的 mock 模块。""" @@ -88,6 +89,7 @@ def _install_mocks(): @classmethod def create(cls, user_id, filename, file_size, storage_key, **kwargs): from uuid import uuid4 + return cls( id=uuid4().hex, user_id=user_id, @@ -111,18 +113,24 @@ def _install_mocks(): # packages.domain, packages.adapters, packages.application namespace for name in [ - "packages", "packages.domain", "packages.ports", - "packages.adapters", "packages.adapters.sqlalchemy_impl", + "packages", + "packages.domain", + "packages.ports", + "packages.adapters", + "packages.adapters.sqlalchemy_impl", "packages.adapters.sqlalchemy_impl.user_repository", "packages.adapters.sqlalchemy_impl.duplication_repository", "packages.adapters.sqlalchemy_impl.session", - "packages.adapters.redis", "packages.adapters.smtp", + "packages.adapters.redis", + "packages.adapters.smtp", ]: if name not in sys.modules: sys.modules[name] = types.ModuleType(name) sys.modules["packages.adapters.sqlalchemy_impl.user_repository"].SQLAlchemyUserRepository = MagicMock - sys.modules["packages.adapters.sqlalchemy_impl.duplication_repository"].SQLAlchemyDuplicationRecordRepository = MagicMock + sys.modules["packages.adapters.sqlalchemy_impl.duplication_repository"].SQLAlchemyDuplicationRecordRepository = ( + MagicMock + ) sys.modules["packages.adapters.sqlalchemy_impl.session"].build_session_factory = MagicMock( return_value=(MagicMock(), MagicMock()) ) @@ -146,6 +154,7 @@ def _install_mocks(): class UploadForDuplicationUseCase: def __init__(self, repo): self.repo = repo + def execute(self, cmd): record = DuplicationRecord.create( user_id=cmd.user_id, @@ -156,20 +165,32 @@ def _install_mocks(): return record class ListDuplicationRecordsUseCase: - def __init__(self, repo): self.repo = repo - def execute(self, user_id, **kw): return [] + def __init__(self, repo): + self.repo = repo + + def execute(self, user_id, **kw): + return [] class GetDuplicationDetailUseCase: - def __init__(self, repo): self.repo = repo - def execute(self, record_id): return None + def __init__(self, repo): + self.repo = repo + + def execute(self, record_id): + return None class DeleteDuplicationRecordUseCase: - def __init__(self, repo): self.repo = repo - def execute(self, record_id): return True + def __init__(self, repo): + self.repo = repo + + def execute(self, record_id): + return True class RetryDuplicationUseCase: - def __init__(self, repo): self.repo = repo - def execute(self, record_id): return None + def __init__(self, repo): + self.repo = repo + + def execute(self, record_id): + return None app_mod.UploadForDuplicationCommand = UploadForDuplicationCommand app_mod.UploadForDuplicationUseCase = UploadForDuplicationUseCase @@ -300,9 +321,8 @@ for ns in ["app", "app.api", "app.api.routes"]: sys.modules[ns] = types.ModuleType(ns) import importlib.util -_spec = importlib.util.spec_from_file_location( - "app.api.routes.duplication", "/tmp/duplication_routes_fixed.py" -) + +_spec = importlib.util.spec_from_file_location("app.api.routes.duplication", "/tmp/duplication_routes_fixed.py") duplication = importlib.util.module_from_spec(_spec) sys.modules["app.api.routes.duplication"] = duplication _spec.loader.exec_module(duplication) @@ -312,6 +332,7 @@ _spec.loader.exec_module(duplication) # 2. Fixtures # --------------------------------------------------------------------------- + def _make_user(**overrides) -> User: defaults = dict( id="user-dup-001", @@ -330,15 +351,26 @@ def _make_user(**overrides) -> User: class MockDuplicationRepo: """内存中的查重记录 Repository mock。""" - def create(self, record): return record - def get(self, record_id): return None - def list_by_user(self, user_id, **kw): return [] - def update(self, record): return record - def delete(self, record_id): return True + + def create(self, record): + return record + + def get(self, record_id): + return None + + def list_by_user(self, user_id, **kw): + return [] + + def update(self, record): + return record + + def delete(self, record_id): + return True class MockStorageService: """可控的存储服务 mock。""" + def __init__(self, should_fail=False, error_msg="Internal server error details"): self.should_fail = should_fail self.error_msg = error_msg @@ -386,6 +418,7 @@ def client(mock_dup_repo, mock_storage): # 3. MIME 类型验证(P0 修复验证) # --------------------------------------------------------------------------- + class TestMIMETypeValidation: """验证 MIME 类型白名单校验。""" @@ -519,6 +552,7 @@ class TestMIMETypeValidation: # 4. 文件大小限制(P0 修复验证) # --------------------------------------------------------------------------- + class TestFileSizeLimit: """验证文件大小限制。""" @@ -555,6 +589,7 @@ class TestFileSizeLimit: # 5. 错误信息不泄露内部异常(P1 核心修复验证) # --------------------------------------------------------------------------- + class TestErrorInfoLeakPrevention: """P1 修复核心:验证错误响应不泄露内部异常堆栈和详细信息。""" @@ -628,7 +663,7 @@ class TestErrorInfoLeakPrevention: assert resp.status_code == 415 body = resp.text assert "Traceback" not in body - assert "File \"" not in body + assert 'File "' not in body assert "line " not in body def test_error_response_no_internal_paths(self, client): @@ -672,6 +707,7 @@ class TestErrorInfoLeakPrevention: # 6. 正常上传流程(验证修复不影响正常功能) # --------------------------------------------------------------------------- + class TestNormalUploadFlow: """验证正常上传流程不受修复影响。""" @@ -738,6 +774,7 @@ class TestNormalUploadFlow: # 7. 边界情况 # --------------------------------------------------------------------------- + class TestEdgeCases: def test_missing_filename_returns_400(self, client): @@ -765,6 +802,7 @@ class TestEdgeCases: # 8. _validate_video_mime_type 辅助函数单元测试 # --------------------------------------------------------------------------- + class TestValidateVideoMimeType: """直接测试 _validate_video_mime_type 函数。""" @@ -785,8 +823,13 @@ class TestValidateVideoMimeType: def test_all_allowed_types_pass(self): """所有允许的 MIME 类型都应通过。""" allowed = [ - "video/mp4", "video/mpeg", "video/quicktime", "video/x-msvideo", - "video/webm", "video/x-matroska", "video/3gpp", + "video/mp4", + "video/mpeg", + "video/quicktime", + "video/x-msvideo", + "video/webm", + "video/x-matroska", + "video/3gpp", ] for mime in allowed: result = duplication._validate_video_mime_type(mime) @@ -795,6 +838,7 @@ class TestValidateVideoMimeType: def test_empty_content_type_raises_400(self): """空 Content-Type 应抛出 400。""" from fastapi import HTTPException + with pytest.raises(HTTPException) as exc_info: duplication._validate_video_mime_type("") # 空字符串 split 后为空,不在白名单 → 415 @@ -805,6 +849,7 @@ class TestValidateVideoMimeType: def test_none_content_type_raises_400(self): """None Content-Type 应抛出 400。""" from fastapi import HTTPException + with pytest.raises(HTTPException) as exc_info: duplication._validate_video_mime_type(None) assert exc_info.value.status_code == 400 @@ -812,6 +857,7 @@ class TestValidateVideoMimeType: def test_invalid_mime_raises_415(self): """无效 MIME 类型应抛出 415。""" from fastapi import HTTPException + with pytest.raises(HTTPException) as exc_info: duplication._validate_video_mime_type("text/html") assert exc_info.value.status_code == 415 @@ -819,6 +865,7 @@ class TestValidateVideoMimeType: def test_415_message_is_safe(self): """415 错误消息不包含技术实现细节。""" from fastapi import HTTPException + with pytest.raises(HTTPException) as exc_info: duplication._validate_video_mime_type("application/json") detail = exc_info.value.detail diff --git a/tests/integration/test_generation_pipeline.py b/tests/integration/test_generation_pipeline.py index 06ae73724..6a18166f8 100644 --- a/tests/integration/test_generation_pipeline.py +++ b/tests/integration/test_generation_pipeline.py @@ -59,9 +59,7 @@ def simulate_generate_video( task.started_at = task.started_at or datetime.now(timezone.utc) task_repo.update(task) - file_url = ( - f"/projects/{task.project_id}/generated/{task.id}/{task.id}.mp4" - ) + file_url = f"/projects/{task.project_id}/generated/{task.id}/{task.id}.mp4" video = GeneratedVideo.create( project_id=task.project_id, generation_task_id=task.id, diff --git a/tests/integration/test_projects.py b/tests/integration/test_projects.py index 993e4d6d4..e4df3f401 100644 --- a/tests/integration/test_projects.py +++ b/tests/integration/test_projects.py @@ -38,7 +38,6 @@ def test_get_project_by_id_restores_workspace_context(): create_use_case = CreateProjectUseCase(repository) get_use_case = GetProjectUseCase(repository) - retrieved = get_use_case.execute(project.id) assert retrieved is not None assert retrieved.id == project.id diff --git a/tests/integration/test_subscription_api.py b/tests/integration/test_subscription_api.py index 048c22ade..a0b80eda4 100644 --- a/tests/integration/test_subscription_api.py +++ b/tests/integration/test_subscription_api.py @@ -10,6 +10,7 @@ 测试使用 FastAPI TestClient + 依赖覆盖(dependency_overrides), 不连接真实数据库,不访问外部服务。 """ + from __future__ import annotations import sys @@ -23,11 +24,11 @@ import pytest from fastapi import FastAPI from fastapi.testclient import TestClient - # --------------------------------------------------------------------------- # 1. Mock 项目内部模块(使 subscription 路由可独立导入) # --------------------------------------------------------------------------- + def _install_mocks(): """在 sys.modules 中安装所有必需的 mock 模块,使 subscription.py 可导入。""" @@ -58,24 +59,41 @@ def _install_mocks(): # ---------- packages.ports.user_repository ---------- class UserRepository: - def save(self, user): pass - def find_by_id(self, user_id): return None - def find_by_email(self, email): return None - def find_by_username(self, username): return None - def find_by_verification_token(self, token): return None - def find_by_password_reset_token(self, token): return None - def delete(self, user_id): return True + def save(self, user): + pass + + def find_by_id(self, user_id): + return None + + def find_by_email(self, email): + return None + + def find_by_username(self, username): + return None + + def find_by_verification_token(self, token): + return None + + def find_by_password_reset_token(self, token): + return None + + def delete(self, user_id): + return True user_repo_mod = types.ModuleType("packages.ports.user_repository") user_repo_mod.UserRepository = UserRepository # ---------- packages (namespace) ---------- for name in [ - "packages", "packages.domain", "packages.ports", - "packages.adapters", "packages.adapters.sqlalchemy_impl", + "packages", + "packages.domain", + "packages.ports", + "packages.adapters", + "packages.adapters.sqlalchemy_impl", "packages.adapters.sqlalchemy_impl.user_repository", "packages.adapters.sqlalchemy_impl.session", - "packages.adapters.redis", "packages.adapters.smtp", + "packages.adapters.redis", + "packages.adapters.smtp", "packages.application", ]: if name not in sys.modules: @@ -95,11 +113,16 @@ def _install_mocks(): # Stub 其他 repository ports(dependencies.py 会 import 它们) for port_name in [ - "asset_repository", "asset_library_repository", - "classification_job_repository", "duplication_repository", - "generated_video_repository", "generation_task_repository", - "title_library_repository", "voice_library_repository", - "ingest_job_repository", "project_repository", + "asset_repository", + "asset_library_repository", + "classification_job_repository", + "duplication_repository", + "generated_video_repository", + "generation_task_repository", + "title_library_repository", + "voice_library_repository", + "ingest_job_repository", + "project_repository", ]: mod = types.ModuleType(f"packages.ports.{port_name}") # 动态创建一个 Mock repository class @@ -251,9 +274,8 @@ for ns in ["app", "app.api", "app.api.routes"]: # 导入 subscription 路由 import importlib.util -_spec = importlib.util.spec_from_file_location( - "app.api.routes.subscription", "/tmp/subscription_routes.py" -) + +_spec = importlib.util.spec_from_file_location("app.api.routes.subscription", "/tmp/subscription_routes.py") subscription = importlib.util.module_from_spec(_spec) sys.modules["app.api.routes.subscription"] = subscription _spec.loader.exec_module(subscription) @@ -263,6 +285,7 @@ _spec.loader.exec_module(subscription) # 2. Fixtures # --------------------------------------------------------------------------- + def _make_user(**overrides) -> User: """创建测试用 User 实例。""" defaults = dict( @@ -324,12 +347,14 @@ def pro_client(mock_user_repo): app.include_router(subscription.router) def _override_get_current_user(): - return AuthenticatedUser(user=_make_user( - subscription_plan="pro", - subscription_status="active", - max_projects=-1, - max_storage_gb=100, - )) + return AuthenticatedUser( + user=_make_user( + subscription_plan="pro", + subscription_status="active", + max_projects=-1, + max_storage_gb=100, + ) + ) def _override_get_user_repo(): return mock_user_repo @@ -344,6 +369,7 @@ def pro_client(mock_user_repo): # 3. GET /current — 获取当前订阅信息 # --------------------------------------------------------------------------- + class TestGetCurrentSubscription: """GET /current 端点测试。""" @@ -391,6 +417,7 @@ class TestGetCurrentSubscription: # 4. GET /billing-records — 获取账单记录 # --------------------------------------------------------------------------- + class TestGetBillingRecords: def test_returns_empty_list(self, client): @@ -406,14 +433,18 @@ class TestGetBillingRecords: # 5. POST /change-plan — 变更套餐 # --------------------------------------------------------------------------- + class TestChangePlan: def test_upgrade_free_to_standard(self, client, mock_user_repo): """从 free 升级到 standard 应成功。""" - resp = client.post("/change-plan", json={ - "target_plan_id": "standard", - "billing_cycle": "monthly", - }) + resp = client.post( + "/change-plan", + json={ + "target_plan_id": "standard", + "billing_cycle": "monthly", + }, + ) assert resp.status_code == 200 data = resp.json() assert data["success"] is True @@ -424,10 +455,13 @@ class TestChangePlan: def test_upgrade_free_to_pro(self, client, mock_user_repo): """从 free 升级到 pro 应成功,配额正确更新。""" - resp = client.post("/change-plan", json={ - "target_plan_id": "pro", - "billing_cycle": "yearly", - }) + resp = client.post( + "/change-plan", + json={ + "target_plan_id": "pro", + "billing_cycle": "yearly", + }, + ) assert resp.status_code == 200 data = resp.json() assert data["success"] is True @@ -444,10 +478,13 @@ class TestChangePlan: def test_upgrade_to_enterprise(self, client, mock_user_repo): """升级到 enterprise 套餐。""" - resp = client.post("/change-plan", json={ - "target_plan_id": "enterprise", - "billing_cycle": "monthly", - }) + resp = client.post( + "/change-plan", + json={ + "target_plan_id": "enterprise", + "billing_cycle": "monthly", + }, + ) assert resp.status_code == 200 data = resp.json() assert data["success"] is True @@ -459,10 +496,13 @@ class TestChangePlan: def test_same_plan_returns_failure(self, client): """当前套餐与目标套餐相同时应返回 success=False。""" - resp = client.post("/change-plan", json={ - "target_plan_id": "free", - "billing_cycle": "monthly", - }) + resp = client.post( + "/change-plan", + json={ + "target_plan_id": "free", + "billing_cycle": "monthly", + }, + ) assert resp.status_code == 200 data = resp.json() assert data["success"] is False @@ -470,19 +510,25 @@ class TestChangePlan: def test_invalid_plan_id_returns_400(self, client): """无效套餐 ID 应返回 400。""" - resp = client.post("/change-plan", json={ - "target_plan_id": "ultra_mega_plan", - "billing_cycle": "monthly", - }) + resp = client.post( + "/change-plan", + json={ + "target_plan_id": "ultra_mega_plan", + "billing_cycle": "monthly", + }, + ) assert resp.status_code == 400 assert "无效的套餐ID" in resp.json()["detail"] def test_invalid_billing_cycle_returns_400(self, client): """无效计费周期应返回 400。""" - resp = client.post("/change-plan", json={ - "target_plan_id": "pro", - "billing_cycle": "weekly", - }) + resp = client.post( + "/change-plan", + json={ + "target_plan_id": "pro", + "billing_cycle": "weekly", + }, + ) assert resp.status_code == 400 assert "无效的计费周期" in resp.json()["detail"] @@ -510,10 +556,13 @@ class TestChangePlan: app.dependency_overrides[subscription.get_user_repository] = lambda: mock_user_repo tc = TestClient(app) - resp = tc.post("/change-plan", json={ - "target_plan_id": "standard", - "billing_cycle": "monthly", - }) + resp = tc.post( + "/change-plan", + json={ + "target_plan_id": "standard", + "billing_cycle": "monthly", + }, + ) assert resp.status_code == 200 # 原始 user 对象不变 assert original_user.subscription_plan == "free" @@ -525,6 +574,7 @@ class TestChangePlan: # 6. POST /cancel — 取消订阅 # --------------------------------------------------------------------------- + class TestCancelSubscription: def test_cancel_pro_subscription(self, pro_client, mock_user_repo): @@ -569,6 +619,7 @@ class TestCancelSubscription: # 7. POST /toggle-auto-renew — 切换自动续费 # --------------------------------------------------------------------------- + class TestToggleAutoRenew: def test_enable_auto_renew(self, client): @@ -594,7 +645,7 @@ class TestToggleAutoRenew: def test_invalid_type_returns_422(self, client): """enabled 传非布尔值应返回 422。""" - resp = client.post("/toggle-auto-renew", json={"enabled": [1,2,3]}) + resp = client.post("/toggle-auto-renew", json={"enabled": [1, 2, 3]}) assert resp.status_code == 422 @@ -602,6 +653,7 @@ class TestToggleAutoRenew: # 8. 辅助函数 / 工具测试 # --------------------------------------------------------------------------- + class TestHelperFunctions: def test_get_plan_name_known_plans(self): diff --git a/tests/unit/test_asset_diagnosis.py b/tests/unit/test_asset_diagnosis.py index 4f6e6795c..be73f4bc4 100644 --- a/tests/unit/test_asset_diagnosis.py +++ b/tests/unit/test_asset_diagnosis.py @@ -176,9 +176,7 @@ class TestAssetDiagnosisRoute: def test_find_by_project_called_with_correct_project_id(self): project = Project(id="proj-123", name="Test", owner_user_id="user-1") - library = AssetLibrary( - id="lib-1", name="Lib", project_id="proj-123", kind=AssetLibraryKind.VIDEO - ) + library = AssetLibrary(id="lib-1", name="Lib", project_id="proj-123", kind=AssetLibraryKind.VIDEO) project_repo = _StubProjectRepository({"proj-123": project}) library_repo = _StubAssetLibraryRepository({"lib-1": library}) asset_repo = _StubAssetRepository() diff --git a/tests/unit/test_auth_middleware.py b/tests/unit/test_auth_middleware.py index 2209c5626..0bbebd46f 100644 --- a/tests/unit/test_auth_middleware.py +++ b/tests/unit/test_auth_middleware.py @@ -83,7 +83,6 @@ def test_legacy_middleware_optional_user_returns_user_with_valid_credentials(): def test_workspace_dependency_allows_member_access(): repo = _WorkspaceMemberRepositoryStub(role="member") - assert role == "member" @@ -105,7 +104,6 @@ class _WorkspaceMemberRepositoryStub: from packages.domain.entities import WorkspaceMember - async def _authenticated_user(): from app.auth import get_current_user as get_authenticated_user diff --git a/tests/unit/test_chunked_upload_routes.py b/tests/unit/test_chunked_upload_routes.py index dcbe489c9..4c48bcd25 100644 --- a/tests/unit/test_chunked_upload_routes.py +++ b/tests/unit/test_chunked_upload_routes.py @@ -8,6 +8,7 @@ chunked_upload.py 路由单元测试 - 文件大小校验 - OSS 凭证校验 """ + from __future__ import annotations import os @@ -29,7 +30,6 @@ from fastapi.testclient import TestClient from packages.domain import AssetLibrary, AssetLibraryKind, Project - # --------------------------------------------------------------------------- # Stub 实现(不继承 Port ABC) # --------------------------------------------------------------------------- diff --git a/tests/unit/test_config_oss.py b/tests/unit/test_config_oss.py index 24e3b5263..077c84178 100644 --- a/tests/unit/test_config_oss.py +++ b/tests/unit/test_config_oss.py @@ -7,6 +7,7 @@ config.py OSS 配置字段单元测试 - 字段名与代码引用一致 - pydantic_settings 加载行为 """ + from __future__ import annotations import importlib.util diff --git a/tests/unit/test_form_upload_routes.py b/tests/unit/test_form_upload_routes.py index dcfa631ae..e3b8c6161 100644 --- a/tests/unit/test_form_upload_routes.py +++ b/tests/unit/test_form_upload_routes.py @@ -8,6 +8,7 @@ upload.py 表单上传端点单元测试 - 多格式支持(3):多种视频(4种)/音频(5种)/图片(6种)格式 - MIME验证(6):有效类型、空类型(400)、不支持类型(415) """ + from __future__ import annotations import io @@ -30,7 +31,6 @@ from fastapi.testclient import TestClient from packages.domain import AssetLibrary, AssetLibraryKind, Project - # --------------------------------------------------------------------------- # 测试用 Stub(不继承 Port ABC,因为 Port 定义 async 方法,路由实际使用同步 duck-type) # --------------------------------------------------------------------------- diff --git a/tests/unit/test_recipe_use_cases.py b/tests/unit/test_recipe_use_cases.py index ac6aeacea..43b2de396 100644 --- a/tests/unit/test_recipe_use_cases.py +++ b/tests/unit/test_recipe_use_cases.py @@ -1,4 +1,5 @@ """Recipe use cases unit tests.""" + from __future__ import annotations from datetime import datetime, timezone diff --git a/tests/unit/test_template_use_cases.py b/tests/unit/test_template_use_cases.py index 4585ffbe6..aff04d284 100644 --- a/tests/unit/test_template_use_cases.py +++ b/tests/unit/test_template_use_cases.py @@ -1,6 +1,7 @@ """ Template Use Cases 单元测试 — 剪辑计划模板 CRUD + 业务规则校验 """ + from unittest.mock import MagicMock, Mock import pytest @@ -215,14 +216,18 @@ class TestValidateTemplateUseCase: def test_one_take_with_one_segment_ok(self, use_case, repo): """一镜到底 + 恰好 1 个片段 → 通过.""" seg = TemplateSegment( - id="seg-001", template_id="tmpl-001", segment_order=1, - duration_min=0, duration_max=60, + id="seg-001", + template_id="tmpl-001", + segment_order=1, + duration_min=0, + duration_max=60, ) template = _make_template(mode="one_take", segments=[seg]) repo.get.return_value = template command = ValidateTemplateCommand( - template_id="tmpl-001", user_id="user-001", + template_id="tmpl-001", + user_id="user-001", ) result = use_case.execute(command) @@ -232,15 +237,15 @@ class TestValidateTemplateUseCase: def test_one_take_with_two_segments_raises(self, use_case, repo): """一镜到底 + 2 个片段 → ValidationError.""" segs = [ - TemplateSegment(id=f"seg-{i}", template_id="tmpl-001", segment_order=i, - duration_min=0, duration_max=30) + TemplateSegment(id=f"seg-{i}", template_id="tmpl-001", segment_order=i, duration_min=0, duration_max=30) for i in (1, 2) ] template = _make_template(mode="one_take", segments=segs) repo.get.return_value = template command = ValidateTemplateCommand( - template_id="tmpl-001", user_id="user-001", + template_id="tmpl-001", + user_id="user-001", ) with pytest.raises(ValidationError, match="一镜到底模式必须恰好有 1 个片段"): use_case.execute(command) @@ -248,16 +253,29 @@ class TestValidateTemplateUseCase: def test_voice_over_all_segments_have_material_type_ok(self, use_case, repo): """口播+B-roll + 所有片段都有 material_type → 通过.""" segs = [ - TemplateSegment(id="seg-1", template_id="tmpl-001", segment_order=1, - duration_min=5, duration_max=15, material_type="人物"), - TemplateSegment(id="seg-2", template_id="tmpl-001", segment_order=2, - duration_min=10, duration_max=30, material_type="场景"), + TemplateSegment( + id="seg-1", + template_id="tmpl-001", + segment_order=1, + duration_min=5, + duration_max=15, + material_type="人物", + ), + TemplateSegment( + id="seg-2", + template_id="tmpl-001", + segment_order=2, + duration_min=10, + duration_max=30, + material_type="场景", + ), ] template = _make_template(mode="voice_over", segments=segs) repo.get.return_value = template command = ValidateTemplateCommand( - template_id="tmpl-001", user_id="user-001", + template_id="tmpl-001", + user_id="user-001", ) result = use_case.execute(command) assert result.warnings == [] @@ -265,16 +283,29 @@ class TestValidateTemplateUseCase: def test_voice_over_missing_material_type_raises(self, use_case, repo): """口播+B-roll + 某片段缺少 material_type → ValidationError.""" segs = [ - TemplateSegment(id="seg-1", template_id="tmpl-001", segment_order=1, - duration_min=5, duration_max=15, material_type="人物"), - TemplateSegment(id="seg-2", template_id="tmpl-001", segment_order=2, - duration_min=10, duration_max=30, material_type=None), # 缺失 + TemplateSegment( + id="seg-1", + template_id="tmpl-001", + segment_order=1, + duration_min=5, + duration_max=15, + material_type="人物", + ), + TemplateSegment( + id="seg-2", + template_id="tmpl-001", + segment_order=2, + duration_min=10, + duration_max=30, + material_type=None, + ), # 缺失 ] template = _make_template(mode="voice_over", segments=segs) repo.get.return_value = template command = ValidateTemplateCommand( - template_id="tmpl-001", user_id="user-001", + template_id="tmpl-001", + user_id="user-001", ) with pytest.raises(ValidationError, match="material_type"): use_case.execute(command) @@ -285,7 +316,8 @@ class TestValidateTemplateUseCase: repo.get.return_value = template command = ValidateTemplateCommand( - template_id="tmpl-001", user_id="user-001", + template_id="tmpl-001", + user_id="user-001", voiceover_duration=70.0, # 70/60 = 1.167, within ±30% ) result = use_case.execute(command) @@ -297,7 +329,8 @@ class TestValidateTemplateUseCase: repo.get.return_value = template command = ValidateTemplateCommand( - template_id="tmpl-001", user_id="user-001", + template_id="tmpl-001", + user_id="user-001", voiceover_duration=100.0, # 100/60 = 1.667, exceeds +30% ) result = use_case.execute(command) @@ -311,7 +344,8 @@ class TestValidateTemplateUseCase: repo.get.return_value = template command = ValidateTemplateCommand( - template_id="tmpl-001", user_id="user-001", + template_id="tmpl-001", + user_id="user-001", voiceover_duration=30.0, # 30/60 = 0.5, below -30% ) result = use_case.execute(command) @@ -324,7 +358,8 @@ class TestValidateTemplateUseCase: repo.get.return_value = None command = ValidateTemplateCommand( - template_id="nonexistent", user_id="user-001", + template_id="nonexistent", + user_id="user-001", ) with pytest.raises(NotFoundError): use_case.execute(command) diff --git a/tests/unit/test_title_library_use_cases.py b/tests/unit/test_title_library_use_cases.py index 77a181faa..d154c4c4d 100644 --- a/tests/unit/test_title_library_use_cases.py +++ b/tests/unit/test_title_library_use_cases.py @@ -27,11 +27,11 @@ from packages.application.title_library.use_cases import ( ) from packages.domain.title_library import TitleLibraryItem - # --------------------------------------------------------------------------- # Fixtures # --------------------------------------------------------------------------- + @pytest.fixture def mock_repo(): """创建 Mock 仓储""" @@ -89,6 +89,7 @@ def existing_title_item(): # 1. CreateTitleLibraryUseCase 测试 # =========================================================================== + class TestCreateTitleLibraryUseCase: """标题库创建 UseCase 测试""" @@ -136,6 +137,7 @@ class TestCreateTitleLibraryUseCase: # 2. 配额逻辑测试(titles: free=50, basic=500, premium=500) # =========================================================================== + class TestCreateTitleLibraryQuota: """标题库创建配额检查测试""" @@ -276,6 +278,7 @@ class TestCreateTitleLibraryQuota: # 3. UpdateTitleLibraryUseCase 测试 # =========================================================================== + class TestUpdateTitleLibraryUseCase: """标题库更新 UseCase 测试""" @@ -389,6 +392,7 @@ class TestUpdateTitleLibraryUseCase: # 4. DeleteTitleLibraryUseCase 测试 # =========================================================================== + class TestDeleteTitleLibraryUseCase: """标题库删除 UseCase 测试""" @@ -416,6 +420,7 @@ class TestDeleteTitleLibraryUseCase: # 5. GetTitleLibraryUseCase 测试 # =========================================================================== + class TestGetTitleLibraryUseCase: """标题库查询 UseCase 测试""" @@ -450,6 +455,7 @@ class TestGetTitleLibraryUseCase: # 6. ListTitleLibraryUseCase 测试 # =========================================================================== + class TestListTitleLibraryUseCase: """标题库列表 UseCase 测试""" @@ -474,9 +480,7 @@ class TestListTitleLibraryUseCase: use_case.execute("user-001", category="新闻", skip=5, limit=10) - mock_repo.list_by_user.assert_called_once_with( - "user-001", category="新闻", skip=5, limit=10 - ) + mock_repo.list_by_user.assert_called_once_with("user-001", category="新闻", skip=5, limit=10) def test_list_empty(self, mock_repo): """测试空列表""" diff --git a/tests/unit/test_upload_routes.py b/tests/unit/test_upload_routes.py index 785d4331a..b3afc7d4f 100644 --- a/tests/unit/test_upload_routes.py +++ b/tests/unit/test_upload_routes.py @@ -8,6 +8,7 @@ upload.py 路由单元测试 - 文件类型校验 - 异常处理路径 """ + from __future__ import annotations import os @@ -29,7 +30,6 @@ from fastapi.testclient import TestClient from packages.domain import AssetLibrary, AssetLibraryKind, Project - # --------------------------------------------------------------------------- # 测试用 Stub(不继承 Port ABC,因为 Port 定义 async 方法,路由实际使用同步 duck-type) # --------------------------------------------------------------------------- diff --git a/tests/unit/test_voice_library_use_cases.py b/tests/unit/test_voice_library_use_cases.py index 6e1ccfd7d..80a2049dc 100644 --- a/tests/unit/test_voice_library_use_cases.py +++ b/tests/unit/test_voice_library_use_cases.py @@ -27,11 +27,11 @@ from packages.application.voice_library.use_cases import ( ) from packages.domain.voice_library import VoiceLibraryItem - # --------------------------------------------------------------------------- # Fixtures # --------------------------------------------------------------------------- + @pytest.fixture def mock_repo(): """创建 Mock 仓储""" @@ -100,6 +100,7 @@ def existing_voice_item(): # 1. CreateVoiceLibraryUseCase 测试 # =========================================================================== + class TestCreateVoiceLibraryUseCase: """配音库创建 UseCase 测试""" @@ -144,11 +145,9 @@ class TestCreateVoiceLibraryUseCase: result = create_use_case.execute(command, plan_name="free") # 核心断言:voice_id 必须来自 command.voice_id - assert result.voice_id == "specific-voice-id-xyz", \ - "voice_id 应来自 command.voice_id,而非其他字段" + assert result.voice_id == "specific-voice-id-xyz", "voice_id 应来自 command.voice_id,而非其他字段" # 同时确保 item 自身生成的 id 与 voice_id 不同 - assert result.id != "specific-voice-id-xyz", \ - "item.id(UUID)不应与 voice_id 混淆" + assert result.id != "specific-voice-id-xyz", "item.id(UUID)不应与 voice_id 混淆" def test_create_voice_id_empty_string(self, create_use_case, mock_repo): """测试 voice_id 为空字符串的合法场景""" @@ -202,6 +201,7 @@ class TestCreateVoiceLibraryUseCase: # 2. 配额逻辑测试(Create 时的配额检查) # =========================================================================== + class TestCreateVoiceLibraryQuota: """配音库创建配额检查测试""" @@ -342,6 +342,7 @@ class TestCreateVoiceLibraryQuota: # 3. UpdateVoiceLibraryUseCase 测试 # =========================================================================== + class TestUpdateVoiceLibraryUseCase: """配音库更新 UseCase 测试""" @@ -399,8 +400,7 @@ class TestUpdateVoiceLibraryUseCase: result = update_use_case.execute(command) # 核心断言:voice_id 应被更新为新值 - assert result.voice_id == "completely-different-voice-id", \ - "voice_id 应被更新为 command.voice_id 的值" + assert result.voice_id == "completely-different-voice-id", "voice_id 应被更新为 command.voice_id 的值" # item 自身的 id 保持不变 assert result.id == "existing-voice-001" @@ -511,6 +511,7 @@ class TestUpdateVoiceLibraryUseCase: # 4. DeleteVoiceLibraryUseCase 测试 # =========================================================================== + class TestDeleteVoiceLibraryUseCase: """配音库删除 UseCase 测试""" @@ -538,6 +539,7 @@ class TestDeleteVoiceLibraryUseCase: # 5. GetVoiceLibraryUseCase 测试 # =========================================================================== + class TestGetVoiceLibraryUseCase: """配音库查询 UseCase 测试""" @@ -573,6 +575,7 @@ class TestGetVoiceLibraryUseCase: # 6. ListVoiceLibraryUseCase 测试 # =========================================================================== + class TestListVoiceLibraryUseCase: """配音库列表 UseCase 测试""" @@ -597,9 +600,7 @@ class TestListVoiceLibraryUseCase: use_case.execute("user-001", status="completed", skip=10, limit=20) - mock_repo.list_by_user.assert_called_once_with( - "user-001", status="completed", skip=10, limit=20 - ) + mock_repo.list_by_user.assert_called_once_with("user-001", status="completed", skip=10, limit=20) def test_list_empty(self, mock_repo): """测试空列表""" @@ -615,6 +616,7 @@ class TestListVoiceLibraryUseCase: # 7. voice_id 与 id 字段隔离专项回归测试 # =========================================================================== + class TestVoiceIdFieldIsolation: """ PR#74 P0 Bug 回归:voice_id 与 item.id 字段隔离 diff --git a/update_tracker.py b/update_tracker.py index 33447886b..f27e6b9ee 100644 --- a/update_tracker.py +++ b/update_tracker.py @@ -1,7 +1,7 @@ import sqlite3 from datetime import datetime -conn = sqlite3.connect('/app/tracker.db') +conn = sqlite3.connect("/app/tracker.db") cursor = conn.cursor() # 先查看当前 Phase 4 任务 @@ -13,10 +13,30 @@ for name, status in tasks[:5]: # Phase 4 已完成的关键任务 phase4_completed_keywords = [ - 'JWT', 'Password', 'Redis', 'Email', '注册', '登录', '登出', '密码重置', - '工作空间', '邀请', '成员', '权限', '订阅', 'Repository', 'API', - 'Docker', 'Kubernetes', '健康检查', 'Celery', 'GitHub', '测试', - '文档', 'MIT', 'README' + "JWT", + "Password", + "Redis", + "Email", + "注册", + "登录", + "登出", + "密码重置", + "工作空间", + "邀请", + "成员", + "权限", + "订阅", + "Repository", + "API", + "Docker", + "Kubernetes", + "健康检查", + "Celery", + "GitHub", + "测试", + "文档", + "MIT", + "README", ] # 更新所有包含关键词的 Phase 4 任务为已完成 @@ -24,25 +44,30 @@ now = datetime.now().isoformat() updated = 0 for keyword in phase4_completed_keywords: - cursor.execute(''' + cursor.execute( + """ UPDATE tasks SET status = 'completed', updated_at = ? WHERE (phase LIKE "%Phase 4%" OR phase LIKE "%Phase%4%") AND (name LIKE ? OR description LIKE ?) AND status != 'completed' - ''', (now, f'%{keyword}%', f'%{keyword}%')) + """, + (now, f"%{keyword}%", f"%{keyword}%"), + ) updated += cursor.rowcount conn.commit() # 统计结果 -cursor.execute('SELECT COUNT(*) FROM tasks WHERE (phase LIKE "%Phase 4%" OR phase LIKE "%Phase%4%") AND status = "completed"') +cursor.execute( + 'SELECT COUNT(*) FROM tasks WHERE (phase LIKE "%Phase 4%" OR phase LIKE "%Phase%4%") AND status = "completed"' +) completed = cursor.fetchone()[0] cursor.execute('SELECT COUNT(*) FROM tasks WHERE phase LIKE "%Phase 4%" OR phase LIKE "%Phase%4%"') total = cursor.fetchone()[0] -print(f'\n✅ 更新完成:') -print(f' - 本次更新: {updated} 个任务') -print(f' - Phase 4 进度: {completed}/{total} 已完成 ({completed/total*100:.1f}%)') +print(f"\n✅ 更新完成:") +print(f" - 本次更新: {updated} 个任务") +print(f" - Phase 4 进度: {completed}/{total} 已完成 ({completed/total*100:.1f}%)") conn.close()