Merge develop into main - v0.1.123
Auto Merge PRs / auto-merge (push) Failing after 3m27s
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 47h9m34s
CI/CD Pipeline / Frontend Lint (push) Failing after 47h8m47s
CI/CD Pipeline / Deploy Staging (push) Failing after 47h6m52s
CI/CD Pipeline / Deploy Production (push) Failing after 1705h41m20s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 1705h41m20s
CI/CD Pipeline / Build Production Runtime Images (push) Failing after 1705h41m52s
CI/CD Pipeline / Production Browser E2E (push) Failing after 1705h41m18s
Auto Merge PRs / auto-merge (push) Failing after 3m27s
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 47h9m34s
CI/CD Pipeline / Frontend Lint (push) Failing after 47h8m47s
CI/CD Pipeline / Deploy Staging (push) Failing after 47h6m52s
CI/CD Pipeline / Deploy Production (push) Failing after 1705h41m20s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 1705h41m20s
CI/CD Pipeline / Build Production Runtime Images (push) Failing after 1705h41m52s
CI/CD Pipeline / Production Browser E2E (push) Failing after 1705h41m18s
This commit is contained in:
@@ -0,0 +1,68 @@
|
||||
"""Add tags and asset_tags tables
|
||||
|
||||
Revision ID: 030
|
||||
Revises: 029
|
||||
Create Date: 2026-07-07
|
||||
|
||||
新增标签表和素材-标签关联表,支持规范化多对多标签管理。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "030"
|
||||
down_revision = "029"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _table_exists(table: str) -> bool:
|
||||
ctx = op.get_context()
|
||||
if ctx.as_sql:
|
||||
return False
|
||||
conn = op.get_bind()
|
||||
result = conn.execute(
|
||||
sa.text("SELECT COUNT(*) FROM information_schema.tables WHERE table_name = :table"),
|
||||
{"table": table},
|
||||
)
|
||||
return (result.scalar() or 0) > 0
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if not _table_exists("tags"):
|
||||
op.create_table(
|
||||
"tags",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("user_id", sa.String(36), nullable=False),
|
||||
sa.Column("name", sa.String(100), nullable=False),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(),
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
),
|
||||
sa.UniqueConstraint("user_id", "name", name="uq_tags_user_name"),
|
||||
)
|
||||
op.create_index("ix_tags_user_id", "tags", ["user_id"])
|
||||
|
||||
if not _table_exists("asset_tags"):
|
||||
op.create_table(
|
||||
"asset_tags",
|
||||
sa.Column("asset_id", sa.String(36), primary_key=True),
|
||||
sa.Column("tag_id", sa.String(36), primary_key=True),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(),
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
),
|
||||
)
|
||||
op.create_index("ix_asset_tags_tag_id", "asset_tags", ["tag_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_asset_tags_tag_id", table_name="asset_tags")
|
||||
op.drop_table("asset_tags")
|
||||
op.drop_index("ix_tags_user_id", table_name="tags")
|
||||
op.drop_table("tags")
|
||||
@@ -16,6 +16,7 @@ from app.api.routes.jobs import router as jobs_router
|
||||
from app.api.routes.projects import router as projects_router
|
||||
from app.api.routes.recipes import router as recipes_router
|
||||
from app.api.routes.subscription import router as subscription_router
|
||||
from app.api.routes.tags import router as tags_router
|
||||
from app.api.routes.task_center import router as task_center_router
|
||||
from app.api.routes.templates import router as templates_router
|
||||
from app.api.routes.titles import router as titles_router
|
||||
@@ -38,6 +39,11 @@ api_router.include_router(
|
||||
prefix="/projects",
|
||||
tags=["Project"],
|
||||
)
|
||||
api_router.include_router(
|
||||
tags_router,
|
||||
prefix="/tags",
|
||||
tags=["Tag"],
|
||||
)
|
||||
api_router.include_router(
|
||||
task_center_router,
|
||||
tags=["TaskCenter"],
|
||||
|
||||
@@ -7,14 +7,18 @@ from app.dependencies import (
|
||||
get_asset_library_repository,
|
||||
get_asset_repository,
|
||||
get_project_repository,
|
||||
get_tag_repository,
|
||||
)
|
||||
from app.schemas.asset import (
|
||||
AssetResponse,
|
||||
BatchDeleteRequest,
|
||||
BatchDeleteResponse,
|
||||
CreateAssetRequest,
|
||||
ListAssetsResponse,
|
||||
UpdateAssetRequest,
|
||||
UpdateAssetReviewRequest,
|
||||
)
|
||||
from app.schemas.tag import TagAssetsRequest
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
from packages.application import (
|
||||
@@ -64,6 +68,7 @@ def _to_asset_response(item, storage_service=None) -> AssetResponse:
|
||||
classification_status=item.classification_status.value,
|
||||
quality_score=item.quality_score,
|
||||
uploaded_by_user_id=item.uploaded_by_user_id,
|
||||
tag_ids=getattr(item, "tag_ids", []),
|
||||
)
|
||||
|
||||
|
||||
@@ -84,6 +89,7 @@ def list_assets(
|
||||
keyword: Optional[str] = Query(None, description="按名称模糊匹配"),
|
||||
gender: Optional[str] = Query(None, description="按 metadata.gender 筛选"),
|
||||
style: Optional[str] = Query(None, description="按 metadata.style 筛选"),
|
||||
tag_ids: Optional[str] = Query(None, description="按标签 ID 筛选(逗号分隔,取交集)"),
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
@@ -96,12 +102,19 @@ def list_assets(
|
||||
# kind → file_type 映射(voice 对应 audio)
|
||||
kind_to_file_type = {"video": "video", "voice": "audio", "image": "image"}
|
||||
|
||||
def _apply_filters(items):
|
||||
"""依次应用 kind / keyword / gender / style 过滤。"""
|
||||
# 解析 tag_ids 参数(逗号分隔)
|
||||
filter_tag_ids: list[str] | None = None
|
||||
if tag_ids:
|
||||
filter_tag_ids = [t.strip() for t in tag_ids.split(",") if t.strip()]
|
||||
if not filter_tag_ids:
|
||||
filter_tag_ids = None
|
||||
|
||||
# 需要内存过滤的标志(keyword/gender/style/tag_ids 无法在 DB 层过滤)
|
||||
needs_memory_filter = bool(keyword or gender or style or filter_tag_ids)
|
||||
|
||||
def _apply_memory_filters(items):
|
||||
"""应用 keyword / gender / style / tag_ids 内存过滤。"""
|
||||
result = items
|
||||
if kind:
|
||||
ft = kind_to_file_type.get(kind)
|
||||
result = [i for i in result if i.mime_type and i.mime_type.startswith(ft or "")]
|
||||
if keyword:
|
||||
kw = keyword.lower()
|
||||
result = [i for i in result if kw in (i.name or "").lower()]
|
||||
@@ -109,9 +122,90 @@ def list_assets(
|
||||
result = [i for i in result if (i.metadata or {}).get("gender") == gender]
|
||||
if style:
|
||||
result = [i for i in result if (i.metadata or {}).get("style") == style]
|
||||
if filter_tag_ids:
|
||||
tag_set = set(filter_tag_ids)
|
||||
result = [i for i in result if tag_set.issubset(set(getattr(i, "tag_ids", [])))]
|
||||
return result
|
||||
|
||||
# 模式1:指定 library_id → 返回该库的素材
|
||||
# ── 优化路径:无内存过滤时,使用 DB 级分页 ──
|
||||
if not needs_memory_filter:
|
||||
ft = kind_to_file_type.get(kind) if kind else None
|
||||
|
||||
# 模式1:指定 library_id
|
||||
if library_id:
|
||||
library = asset_library_repository.get(library_id)
|
||||
if library is None:
|
||||
raise HTTPException(status_code=404, detail=f"AssetLibrary {library_id} not found")
|
||||
_check_project_access(library.project_id, user_id, project_repository)
|
||||
if ft:
|
||||
items = asset_repository.find_by_library_and_file_type(library_id, ft, skip=skip, limit=limit)
|
||||
total = asset_repository.count_by_project(library.project_id) if not kind else len(items)
|
||||
else:
|
||||
items = asset_repository.find_by_library(library_id, skip=skip, limit=limit)
|
||||
total = asset_repository.count_by_project(library.project_id)
|
||||
return ListAssetsResponse(
|
||||
items=[_to_asset_response(item) for item in items],
|
||||
total=total,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
# 模式2:指定 project_id
|
||||
if project_id:
|
||||
_check_project_access(project_id, user_id, project_repository)
|
||||
if ft:
|
||||
# 无直接方法,加载后按 file_type 过滤(仍比全量加载好)
|
||||
all_items = asset_repository.find_by_project(project_id)
|
||||
items = [i for i in all_items if i.mime_type and i.mime_type.startswith(ft)]
|
||||
total = len(items)
|
||||
paged = items[skip : skip + limit]
|
||||
else:
|
||||
items = asset_repository.find_by_project(project_id, skip=skip, limit=limit)
|
||||
total = asset_repository.count_by_project(project_id)
|
||||
paged = items
|
||||
return ListAssetsResponse(
|
||||
items=[_to_asset_response(item) for item in paged],
|
||||
total=total,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
# 模式3:跨项目(无 library_id/project_id)
|
||||
try:
|
||||
projects = project_repository.find_accessible_projects(user_id)
|
||||
except Exception:
|
||||
logger.exception("查询用户可访问项目失败: user_id=%s", user_id)
|
||||
return ListAssetsResponse(items=[], total=0, skip=skip, limit=limit)
|
||||
|
||||
project_ids = [p.id for p in projects]
|
||||
if not project_ids:
|
||||
return ListAssetsResponse(items=[], total=0, skip=skip, limit=limit)
|
||||
|
||||
total = asset_repository.count_by_project_ids(project_ids)
|
||||
# 跨项目分页:逐项目累积直到凑够一页
|
||||
paged_items: list = []
|
||||
offset = skip
|
||||
remaining = limit
|
||||
for pid in project_ids:
|
||||
proj_total = asset_repository.count_by_project(pid)
|
||||
if offset >= proj_total:
|
||||
offset -= proj_total
|
||||
continue
|
||||
proj_items = asset_repository.find_by_project(pid, skip=offset, limit=remaining)
|
||||
paged_items.extend(proj_items)
|
||||
remaining -= len(proj_items)
|
||||
offset = 0
|
||||
if remaining <= 0:
|
||||
break
|
||||
|
||||
return ListAssetsResponse(
|
||||
items=[_to_asset_response(item) for item in paged_items],
|
||||
total=total,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
# ── 内存过滤路径:有 keyword/gender/style 时,加载全量后内存过滤 ──
|
||||
if library_id:
|
||||
library = asset_library_repository.get(library_id)
|
||||
if library is None:
|
||||
@@ -121,41 +215,24 @@ def list_assets(
|
||||
all_items = asset_repository.find_by_library_and_file_type(library_id, kind_to_file_type[kind])
|
||||
else:
|
||||
all_items = asset_repository.find_by_library(library_id)
|
||||
filtered = _apply_filters(all_items)
|
||||
total = len(filtered)
|
||||
paged = filtered[skip : skip + limit]
|
||||
return ListAssetsResponse(
|
||||
items=[_to_asset_response(item) for item in paged],
|
||||
total=total,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
# 模式2:指定 project_id → 返回该项目所有素材
|
||||
if project_id:
|
||||
elif project_id:
|
||||
_check_project_access(project_id, user_id, project_repository)
|
||||
all_items = asset_repository.find_by_project(project_id)
|
||||
filtered = _apply_filters(all_items)
|
||||
total = len(filtered)
|
||||
paged = filtered[skip : skip + limit]
|
||||
return ListAssetsResponse(
|
||||
items=[_to_asset_response(item) for item in paged],
|
||||
total=total,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
)
|
||||
else:
|
||||
try:
|
||||
projects = project_repository.find_accessible_projects(user_id)
|
||||
except Exception:
|
||||
logger.exception("查询用户可访问项目失败: user_id=%s", user_id)
|
||||
return ListAssetsResponse(items=[], total=0, skip=skip, limit=limit)
|
||||
all_items = []
|
||||
for proj in projects:
|
||||
all_items.extend(asset_repository.find_by_project(proj.id))
|
||||
|
||||
# 模式3:都不传 → 返回用户可访问的所有项目的所有素材
|
||||
try:
|
||||
projects = project_repository.find_accessible_projects(user_id)
|
||||
except Exception:
|
||||
logger.exception("查询用户可访问项目失败: user_id=%s", user_id)
|
||||
return ListAssetsResponse(items=[], total=0, skip=skip, limit=limit)
|
||||
|
||||
all_items = []
|
||||
for proj in projects:
|
||||
all_items.extend(asset_repository.find_by_project(proj.id))
|
||||
filtered = _apply_filters(all_items)
|
||||
# 应用 kind 过滤(如果有)+ keyword/gender/style
|
||||
if kind:
|
||||
ft = kind_to_file_type.get(kind)
|
||||
all_items = [i for i in all_items if i.mime_type and i.mime_type.startswith(ft or "")]
|
||||
filtered = _apply_memory_filters(all_items)
|
||||
total = len(filtered)
|
||||
paged = filtered[skip : skip + limit]
|
||||
return ListAssetsResponse(
|
||||
@@ -191,6 +268,35 @@ def update_asset_review_status(
|
||||
return _to_asset_response(updated)
|
||||
|
||||
|
||||
@router.post("/batch-delete", response_model=BatchDeleteResponse)
|
||||
def batch_delete_assets(
|
||||
request: BatchDeleteRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
asset_repository: Any = Depends(get_asset_repository),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
) -> BatchDeleteResponse:
|
||||
"""批量删除素材(配音素材等),需逐项校验项目权限。"""
|
||||
user_id = authenticated_user.user.id
|
||||
deleted_ids: list[str] = []
|
||||
failed_ids: list[str] = []
|
||||
|
||||
for asset_id in request.ids:
|
||||
item = asset_repository.find_by_id(asset_id)
|
||||
if item is None:
|
||||
failed_ids.append(asset_id)
|
||||
continue
|
||||
try:
|
||||
_check_project_access(item.project_id, user_id, project_repository)
|
||||
deleted_ids.append(asset_id)
|
||||
except HTTPException:
|
||||
failed_ids.append(asset_id)
|
||||
|
||||
if deleted_ids:
|
||||
asset_repository.batch_delete(deleted_ids)
|
||||
|
||||
return BatchDeleteResponse(deleted_count=len(deleted_ids), failed_ids=failed_ids)
|
||||
|
||||
|
||||
@router.get("/{asset_id}", response_model=AssetResponse)
|
||||
def get_asset(
|
||||
asset_id: str,
|
||||
@@ -244,6 +350,48 @@ def delete_asset(
|
||||
asset_repository.delete(asset_id)
|
||||
|
||||
|
||||
@router.post("/{asset_id}/tags", response_model=AssetResponse)
|
||||
def tag_asset(
|
||||
asset_id: str,
|
||||
request: TagAssetsRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
asset_repository: Any = Depends(get_asset_repository),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
tag_repository: Any = Depends(get_tag_repository),
|
||||
) -> AssetResponse:
|
||||
"""给素材打标签。"""
|
||||
item = asset_repository.find_by_id(asset_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail=f"Asset {asset_id} not found")
|
||||
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
|
||||
for tag_id in request.tag_ids:
|
||||
tag = tag_repository.get(tag_id)
|
||||
if tag is None:
|
||||
raise HTTPException(status_code=404, detail=f"Tag {tag_id} not found")
|
||||
if tag.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail=f"无权使用标签 {tag_id}")
|
||||
item.add_tag(tag_id)
|
||||
updated = asset_repository.update(item)
|
||||
return _to_asset_response(updated)
|
||||
|
||||
|
||||
@router.delete("/{asset_id}/tags/{tag_id}", status_code=204)
|
||||
def untag_asset(
|
||||
asset_id: str,
|
||||
tag_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
asset_repository: Any = Depends(get_asset_repository),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
) -> None:
|
||||
"""取消素材的标签。"""
|
||||
item = asset_repository.find_by_id(asset_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=404, detail=f"Asset {asset_id} not found")
|
||||
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
|
||||
item.remove_tag(tag_id)
|
||||
asset_repository.update(item)
|
||||
|
||||
|
||||
@router.post("", response_model=AssetResponse)
|
||||
def create_asset(
|
||||
request: CreateAssetRequest,
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
"""标签 CRUD 路由。"""
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_tag_repository
|
||||
from app.schemas.tag import (
|
||||
CreateTagRequest,
|
||||
ListTagsResponse,
|
||||
TagResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from packages.domain import Tag
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get("", response_model=ListTagsResponse)
|
||||
def list_tags(
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
tag_repository: Any = Depends(get_tag_repository),
|
||||
) -> ListTagsResponse:
|
||||
"""列出当前用户的标签。"""
|
||||
user_id = authenticated_user.user.id
|
||||
items = tag_repository.list_by_user(user_id, skip=skip, limit=limit)
|
||||
total = tag_repository.count_by_user(user_id)
|
||||
return ListTagsResponse(
|
||||
items=[TagResponse(id=t.id, name=t.name, created_at=t.created_at) for t in items],
|
||||
total=total,
|
||||
)
|
||||
|
||||
|
||||
@router.post("", response_model=TagResponse, status_code=201)
|
||||
def create_tag(
|
||||
request: CreateTagRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
tag_repository: Any = Depends(get_tag_repository),
|
||||
) -> TagResponse:
|
||||
"""创建标签(同用户同名去重,返回 409)。"""
|
||||
user_id = authenticated_user.user.id
|
||||
existing = tag_repository.find_by_name(user_id, request.name)
|
||||
if existing:
|
||||
raise HTTPException(status_code=409, detail="标签名称已存在")
|
||||
tag = Tag.create(user_id=user_id, name=request.name)
|
||||
created = tag_repository.create(tag)
|
||||
return TagResponse(id=created.id, name=created.name, created_at=created.created_at)
|
||||
|
||||
|
||||
@router.delete("/{tag_id}", status_code=204)
|
||||
def delete_tag(
|
||||
tag_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
tag_repository: Any = Depends(get_tag_repository),
|
||||
) -> None:
|
||||
"""删除标签(同时清理素材关联)。"""
|
||||
tag = tag_repository.get(tag_id)
|
||||
if tag is None:
|
||||
raise HTTPException(status_code=404, detail="标签不存在")
|
||||
if tag.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="无权删除该标签")
|
||||
tag_repository.delete(tag_id)
|
||||
+153
-10
@@ -9,22 +9,28 @@ from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import (
|
||||
get_cosyvoice_service,
|
||||
get_db_session,
|
||||
get_user_repository,
|
||||
get_voice_clone_profile_repository,
|
||||
get_voice_library_repository,
|
||||
)
|
||||
from app.schemas.tts import (
|
||||
ListTTSJobResponse,
|
||||
SaveToLibraryRequest,
|
||||
SaveToLibraryResponse,
|
||||
TTSJobResponse,
|
||||
TTSStatusResponse,
|
||||
TTSSynthesizeRequest,
|
||||
TTSSynthesizeResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Response, WebSocket, WebSocketDisconnect, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.tts_job_repository import (
|
||||
SQLAlchemyTTSJobRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.voice_library_repository import SQLAlchemyVoiceLibraryRepository
|
||||
from packages.application.cosyvoice_service import CosyVoiceService
|
||||
from packages.application.tts_job.streaming_service import TTSStreamingService
|
||||
from packages.application.tts_job.use_cases import (
|
||||
CreateTTSJobUseCase,
|
||||
DeleteTTSJobUseCase,
|
||||
@@ -34,6 +40,12 @@ from packages.application.tts_job.use_cases import (
|
||||
TTSJobNotFoundError,
|
||||
)
|
||||
from packages.application.tts_job.workflow import TTSWorkflowService
|
||||
from packages.application.voice_library.commands import CreateVoiceLibraryCommand
|
||||
from packages.application.voice_library.use_cases import (
|
||||
CreateVoiceLibraryUseCase,
|
||||
QuotaExceededError,
|
||||
)
|
||||
from packages.ports.user_repository import UserRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -130,18 +142,25 @@ def synthesize(
|
||||
|
||||
# 若任务处于 processing 状态(异步模式),触发 Celery 后台轮询
|
||||
if job.status.value == "processing":
|
||||
task_id = (job.metadata or {}).get("cosyvoice_task_id", "")
|
||||
if task_id:
|
||||
try:
|
||||
# 分段合成任务 vs 普通单段任务
|
||||
segment_task_ids = (job.metadata or {}).get("segment_task_ids", [])
|
||||
is_segment = len(segment_task_ids) > 0
|
||||
|
||||
try:
|
||||
if is_segment:
|
||||
from worker_app.tasks import process_tts_segment_synthesis
|
||||
|
||||
process_tts_segment_synthesis.delay(job.id)
|
||||
else:
|
||||
from worker_app.tasks import process_tts_synthesis
|
||||
|
||||
process_tts_synthesis.delay(job.id)
|
||||
except Exception as e:
|
||||
# Celery 调度失败,标记 job 为 failed
|
||||
try:
|
||||
workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}")
|
||||
except Exception as inner_e:
|
||||
logger.error(f"Celery 调度后标记失败时出错: job_id={job.id}, error={inner_e}")
|
||||
except Exception as e:
|
||||
# Celery 调度失败,标记 job 为 failed
|
||||
try:
|
||||
workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}")
|
||||
except Exception as inner_e:
|
||||
logger.error(f"Celery 调度后标记失败时出错: job_id={job.id}, error={inner_e}")
|
||||
|
||||
return TTSSynthesizeResponse(
|
||||
job_id=job.id,
|
||||
@@ -225,3 +244,127 @@ def delete_tts_job(
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="TTS job not found")
|
||||
return Response(status_code=204)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/jobs/{job_id}/save-to-library",
|
||||
response_model=SaveToLibraryResponse,
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
)
|
||||
def save_tts_job_to_library(
|
||||
job_id: str,
|
||||
request: SaveToLibraryRequest = SaveToLibraryRequest(),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
tts_repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
|
||||
voice_library_repository: SQLAlchemyVoiceLibraryRepository = Depends(get_voice_library_repository),
|
||||
user_repository: UserRepository = Depends(get_user_repository),
|
||||
) -> SaveToLibraryResponse:
|
||||
"""将已完成的 TTS 合成结果保存到配音库。
|
||||
|
||||
自动携带音色名、时长、语速等元信息。
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
|
||||
# 获取 TTS job
|
||||
get_use_case = GetTTSJobUseCase(tts_repository)
|
||||
try:
|
||||
job = get_use_case.execute(job_id, user_id)
|
||||
except TTSJobNotFoundError:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="TTS job not found")
|
||||
|
||||
# 校验已完成
|
||||
if not job.is_completed:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="TTS job is not completed yet",
|
||||
)
|
||||
|
||||
# 构建配音素材名称
|
||||
name = request.name or f"TTS-{job.id[:8]}"
|
||||
|
||||
# 构建元信息
|
||||
metadata_ = {
|
||||
"source": "tts_job",
|
||||
"tts_job_id": job.id,
|
||||
"format": job.format,
|
||||
"sample_rate": job.sample_rate,
|
||||
}
|
||||
if job.metadata:
|
||||
# 保留原始 job 的有用元信息
|
||||
for key in ("speed", "language"):
|
||||
if key in job.metadata:
|
||||
metadata_[key] = job.metadata[key]
|
||||
|
||||
# 获取用户套餐(用于配额检查)
|
||||
user = user_repository.find_by_id(user_id)
|
||||
plan_name = getattr(user, "subscription_plan", "free") if user else "free"
|
||||
|
||||
# 构建命令并执行
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id=user_id,
|
||||
name=name,
|
||||
text=job.input_text,
|
||||
voice_provider="cosyvoice",
|
||||
voice_id=job.voice_id,
|
||||
voice_name=job.voice_model or "",
|
||||
audio_url=job.output_audio_url,
|
||||
duration=job.duration,
|
||||
file_size=job.file_size,
|
||||
status="completed",
|
||||
project_id=job.project_id or "",
|
||||
tags=[],
|
||||
metadata_=metadata_,
|
||||
)
|
||||
|
||||
use_case = CreateVoiceLibraryUseCase(voice_library_repository)
|
||||
try:
|
||||
item = use_case.execute(command, plan_name=plan_name or "free")
|
||||
except QuotaExceededError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
||||
detail=f"配音库配额已满({exc.used}/{exc.limit}),请升级套餐",
|
||||
)
|
||||
|
||||
return SaveToLibraryResponse(
|
||||
id=item.id,
|
||||
name=item.name,
|
||||
audio_url=item.audio_url,
|
||||
duration=item.duration,
|
||||
voice_id=item.voice_id,
|
||||
voice_name=item.voice_name,
|
||||
status=item.status,
|
||||
)
|
||||
|
||||
|
||||
@router.websocket("/ws/tts/stream")
|
||||
async def tts_websocket_stream(
|
||||
websocket: WebSocket,
|
||||
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
|
||||
) -> None:
|
||||
"""WebSocket 流式 TTS 合成。
|
||||
|
||||
协议:
|
||||
1. 客户端发送 JSON 文本帧: {"text": "...", "voice_id": "...", ...}
|
||||
2. 服务端发送 JSON 状态帧 + 二进制音频帧
|
||||
3. 完成时发送 JSON 结束帧
|
||||
"""
|
||||
await websocket.accept()
|
||||
try:
|
||||
message = await websocket.receive_json()
|
||||
params = {
|
||||
"text": message.get("text", ""),
|
||||
"voice_id": message.get("voice_id", ""),
|
||||
"sample_rate": message.get("sample_rate", 0),
|
||||
"format": message.get("format", "mp3"),
|
||||
"speed": message.get("speed", 1.0),
|
||||
}
|
||||
streaming_service = TTSStreamingService(cosyvoice_service)
|
||||
await streaming_service.synthesize_and_stream(websocket, params)
|
||||
except WebSocketDisconnect:
|
||||
logger.info("WebSocket 客户端断开连接")
|
||||
except Exception as e:
|
||||
logger.error(f"WebSocket 流式合成异常: {e}", exc_info=True)
|
||||
try:
|
||||
await websocket.send_json({"type": "error", "message": f"服务异常: {e}"})
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@@ -39,6 +39,7 @@ from packages.adapters.sqlalchemy_impl.project_repository import (
|
||||
SQLAlchemyProjectRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.session import build_session_factory
|
||||
from packages.adapters.sqlalchemy_impl.tag_repository import SQLAlchemyTagRepository
|
||||
from packages.adapters.sqlalchemy_impl.title_library_repository import (
|
||||
SQLAlchemyTitleLibraryRepository,
|
||||
)
|
||||
@@ -58,6 +59,7 @@ from packages.ports.generation_task_repository import GenerationTaskRepository
|
||||
from packages.ports.ingest_job_repository import IngestJobRepository
|
||||
from packages.ports.job_repository import JobRepository
|
||||
from packages.ports.project_repository import ProjectRepository
|
||||
from packages.ports.tag_repository import TagRepository
|
||||
from packages.ports.title_library_repository import TitleLibraryRepository
|
||||
from packages.ports.user_repository import UserRepository
|
||||
from packages.ports.voice_clone_profile_repository import VoiceCloneProfileRepository
|
||||
@@ -138,6 +140,13 @@ def get_project_repository(
|
||||
return SQLAlchemyProjectRepository(session)
|
||||
|
||||
|
||||
def get_tag_repository(
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> TagRepository:
|
||||
"""Provide the SQLAlchemy tag repository implementation."""
|
||||
return SQLAlchemyTagRepository(session)
|
||||
|
||||
|
||||
def get_user_repository(
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> UserRepository:
|
||||
|
||||
@@ -51,6 +51,20 @@ class AssetResponse(BaseModel):
|
||||
classification_status: str
|
||||
quality_score: float | None = None
|
||||
uploaded_by_user_id: str
|
||||
tag_ids: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class BatchDeleteRequest(BaseModel):
|
||||
"""批量删除请求。"""
|
||||
|
||||
ids: list[str] = Field(..., min_length=1, max_length=100, description="要删除的素材 ID 列表")
|
||||
|
||||
|
||||
class BatchDeleteResponse(BaseModel):
|
||||
"""批量删除响应。"""
|
||||
|
||||
deleted_count: int = Field(..., ge=0, description="实际删除数量")
|
||||
failed_ids: list[str] = Field(default_factory=list, description="删除失败的 ID 列表")
|
||||
|
||||
|
||||
class ListAssetsResponse(BaseModel):
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
"""标签相关 Schema。"""
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class CreateTagRequest(BaseModel):
|
||||
name: str = Field(..., min_length=1, max_length=100)
|
||||
|
||||
|
||||
class TagResponse(BaseModel):
|
||||
id: str
|
||||
name: str
|
||||
created_at: datetime
|
||||
|
||||
|
||||
class ListTagsResponse(BaseModel):
|
||||
items: list[TagResponse]
|
||||
total: int = Field(default=0, ge=0)
|
||||
|
||||
|
||||
class TagAssetsRequest(BaseModel):
|
||||
tag_ids: list[str] = Field(..., min_length=1, max_length=50)
|
||||
@@ -83,3 +83,21 @@ class ListTTSJobResponse(BaseModel):
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
|
||||
|
||||
class SaveToLibraryRequest(BaseModel):
|
||||
"""保存到配音库请求。"""
|
||||
|
||||
name: Optional[str] = Field(None, description="配音素材名称,留空则自动生成")
|
||||
|
||||
|
||||
class SaveToLibraryResponse(BaseModel):
|
||||
"""保存到配音库响应。"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
audio_url: str
|
||||
duration: float
|
||||
voice_id: str
|
||||
voice_name: str
|
||||
status: str
|
||||
|
||||
@@ -19,6 +19,7 @@ export interface AssetItem {
|
||||
status?: string;
|
||||
classification_status?: string | null;
|
||||
quality_score?: number | null;
|
||||
tag_ids?: string[];
|
||||
created_at?: string;
|
||||
}
|
||||
|
||||
@@ -144,12 +145,18 @@ export const getAssets = async (libraryId: string): Promise<AssetItem[]> => {
|
||||
/** 按类型获取素材(如 voice/video/image),支持可选筛选 */
|
||||
export const getAssetsByKind = async (
|
||||
kind: string,
|
||||
filters?: { keyword?: string; gender?: string; style?: string },
|
||||
filters?: {
|
||||
keyword?: string;
|
||||
gender?: string;
|
||||
style?: string;
|
||||
tag_ids?: string[];
|
||||
},
|
||||
): Promise<AssetItem[]> => {
|
||||
const params: Record<string, string> = { kind };
|
||||
if (filters?.keyword) params.keyword = filters.keyword;
|
||||
if (filters?.gender) params.gender = filters.gender;
|
||||
if (filters?.style) params.style = filters.style;
|
||||
if (filters?.tag_ids?.length) params.tag_ids = filters.tag_ids.join(",");
|
||||
const response = await apiClient.get("/assets", { params });
|
||||
return response.data.items || [];
|
||||
};
|
||||
@@ -233,10 +240,11 @@ export const completeDirectUpload = async (data: {
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 直传上传(大文件推荐) */
|
||||
/** 直传上传(大文件推荐),支持可选进度回调 */
|
||||
export const uploadAssetDirect = async (data: {
|
||||
file: File;
|
||||
library_id: string;
|
||||
onProgress?: (percent: number) => void;
|
||||
}): Promise<{ storage_key: string; ingest_job_id: string }> => {
|
||||
// 后端要求 project_id,前端自动获取默认项目
|
||||
const project = await getOrCreateDefaultProject();
|
||||
@@ -255,13 +263,25 @@ export const uploadAssetDirect = async (data: {
|
||||
);
|
||||
directForm.append("file", data.file);
|
||||
|
||||
const uploadResponse = await fetch(prepared.upload_url, {
|
||||
method: prepared.method,
|
||||
body: directForm,
|
||||
// 使用 XMLHttpRequest 以获取上传进度(fetch 不支持)
|
||||
await new Promise<void>((resolve, reject) => {
|
||||
const xhr = new XMLHttpRequest();
|
||||
xhr.open(prepared.method, prepared.upload_url);
|
||||
xhr.upload.onprogress = (e) => {
|
||||
if (e.lengthComputable && data.onProgress) {
|
||||
data.onProgress(Math.round((e.loaded / e.total) * 100));
|
||||
}
|
||||
};
|
||||
xhr.onload = () => {
|
||||
if (xhr.status >= 200 && xhr.status < 300) {
|
||||
resolve();
|
||||
} else {
|
||||
reject(new Error(`OSS direct upload failed: ${xhr.status}`));
|
||||
}
|
||||
};
|
||||
xhr.onerror = () => reject(new Error("OSS direct upload failed"));
|
||||
xhr.send(directForm);
|
||||
});
|
||||
if (!uploadResponse.ok) {
|
||||
throw new Error(`OSS direct upload failed: ${uploadResponse.status}`);
|
||||
}
|
||||
|
||||
return completeDirectUpload({
|
||||
project_id: project.id,
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
/**
|
||||
* 标签 CRUD API
|
||||
* P3 标签体系:对接后端标签表
|
||||
*/
|
||||
import apiClient from "./client";
|
||||
|
||||
export interface TagItem {
|
||||
id: string;
|
||||
name: string;
|
||||
created_at?: string;
|
||||
usage_count?: number;
|
||||
}
|
||||
|
||||
/** 获取当前用户所有标签 */
|
||||
export const getTags = async (): Promise<TagItem[]> => {
|
||||
const response = await apiClient.get("/tags");
|
||||
return response.data.items || [];
|
||||
};
|
||||
|
||||
/** 创建标签(同名返回 409) */
|
||||
export const createTag = async (name: string): Promise<TagItem> => {
|
||||
const response = await apiClient.post("/tags", { name });
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 删除标签(同时清理素材关联) */
|
||||
export const deleteTag = async (tagId: string): Promise<void> => {
|
||||
await apiClient.delete(`/tags/${tagId}`);
|
||||
};
|
||||
|
||||
/** 为素材添加标签(最多 50 个) */
|
||||
export const tagAsset = async (
|
||||
assetId: string,
|
||||
tagIds: string[],
|
||||
): Promise<void> => {
|
||||
if (tagIds.length === 0) return;
|
||||
await apiClient.post(`/assets/${assetId}/tags`, { tag_ids: tagIds });
|
||||
};
|
||||
|
||||
/** 移除素材的某个标签 */
|
||||
export const untagAsset = async (
|
||||
assetId: string,
|
||||
tagId: string,
|
||||
): Promise<void> => {
|
||||
await apiClient.delete(`/assets/${assetId}/tags/${tagId}`);
|
||||
};
|
||||
@@ -122,6 +122,20 @@ export const getTTSJobs = async (
|
||||
return response.data;
|
||||
};
|
||||
|
||||
/** 存为素材请求参数 */
|
||||
export interface SaveTtsToLibraryRequest {
|
||||
name?: string;
|
||||
tag_ids?: string[];
|
||||
}
|
||||
|
||||
/** 将 TTS 合成结果保存到配音素材库 */
|
||||
export const saveTtsToLibrary = async (
|
||||
jobId: string,
|
||||
data?: SaveTtsToLibraryRequest,
|
||||
): Promise<void> => {
|
||||
await apiClient.post(`/tts/jobs/${jobId}/save-to-library`, data ?? {});
|
||||
};
|
||||
|
||||
/** 删除 TTS 任务 */
|
||||
export const deleteTTSJob = async (jobId: string): Promise<void> => {
|
||||
await apiClient.delete(`/tts/jobs/${jobId}`);
|
||||
|
||||
@@ -1968,7 +1968,9 @@
|
||||
border: 1px dashed var(--primary, #6366f1);
|
||||
border-radius: 6px;
|
||||
cursor: pointer;
|
||||
transition: background 0.15s, color 0.15s;
|
||||
transition:
|
||||
background 0.15s,
|
||||
color 0.15s;
|
||||
}
|
||||
|
||||
.ep-voice-upload-btn:hover {
|
||||
@@ -1996,7 +1998,9 @@
|
||||
border: none;
|
||||
border-radius: 4px;
|
||||
cursor: pointer;
|
||||
transition: background 0.15s, color 0.15s;
|
||||
transition:
|
||||
background 0.15s,
|
||||
color 0.15s;
|
||||
}
|
||||
|
||||
.ep-voice-refresh-btn:hover {
|
||||
@@ -2035,7 +2039,9 @@
|
||||
border: 1px solid var(--border-color, #334155);
|
||||
border-radius: 6px;
|
||||
cursor: pointer;
|
||||
transition: background 0.15s, border-color 0.15s;
|
||||
transition:
|
||||
background 0.15s,
|
||||
border-color 0.15s;
|
||||
}
|
||||
|
||||
.ep-voice-preview-btn:hover {
|
||||
|
||||
@@ -15,6 +15,9 @@ import {
|
||||
PauseCircleOutlined,
|
||||
DownloadOutlined,
|
||||
ShareAltOutlined,
|
||||
SaveOutlined,
|
||||
PlusOutlined,
|
||||
CloseOutlined,
|
||||
} from "@ant-design/icons";
|
||||
import type { AssetItem } from "@/api/assets";
|
||||
import { getAssets, getAssetLibraries } from "@/api/assets";
|
||||
@@ -28,9 +31,10 @@ import type { PresetVoiceItem } from "@/api/voices";
|
||||
import { formatDuration } from "@/api/voiceClone";
|
||||
import type { VoiceClone } from "@/api/voiceClone";
|
||||
import CloneModal from "@/components/voice/CloneModal";
|
||||
import { synthesizeSpeech, getTTSJobStatus } from "@/api/tts";
|
||||
import { synthesizeSpeech, getTTSJobStatus, saveTtsToLibrary } from "@/api/tts";
|
||||
import { getTags, createTag } from "@/api/tags";
|
||||
import { useCloneProgress } from "@/hooks/useCloneProgress";
|
||||
import { useSearchParams } from "react-router-dom";
|
||||
import { useSearchParams, useNavigate } from "react-router-dom";
|
||||
import { getEditPlan } from "@/api/editPlans";
|
||||
import "./generate.css";
|
||||
|
||||
@@ -97,6 +101,8 @@ const STEPS = [
|
||||
================================================================ */
|
||||
|
||||
const GeneratePage: React.FC = () => {
|
||||
const navigate = useNavigate();
|
||||
|
||||
/* ── 步骤状态 ── */
|
||||
const [currentStep, setCurrentStep] = useState(1);
|
||||
|
||||
@@ -225,6 +231,23 @@ const GeneratePage: React.FC = () => {
|
||||
const [customAudioUrl, setCustomAudioUrl] = useState<string | null>(null);
|
||||
const [ttsError, setTtsError] = useState<string | null>(null);
|
||||
const [ttsJobId, setTtsJobId] = useState<string | null>(null);
|
||||
/** 合成完成后保留的 job ID,用于"存为素材" */
|
||||
const [completedTtsJobId, setCompletedTtsJobId] = useState<string | null>(
|
||||
null,
|
||||
);
|
||||
|
||||
/* ── 存为素材弹窗状态 ── */
|
||||
const [saveModalOpen, setSaveModalOpen] = useState(false);
|
||||
const [saveName, setSaveName] = useState("");
|
||||
const [saveTagIds, setSaveTagIds] = useState<string[]>([]);
|
||||
const [saveNewTag, setSaveNewTag] = useState("");
|
||||
|
||||
/* ── 标签列表(用于存为素材弹窗) ── */
|
||||
const { data: allTags = [] } = useQuery({
|
||||
queryKey: ["generate-save-tags"],
|
||||
queryFn: getTags,
|
||||
staleTime: 30_000,
|
||||
});
|
||||
|
||||
/* ── 素材数据 API ── */
|
||||
const { data: libraries = [] } = useQuery({
|
||||
@@ -310,6 +333,7 @@ const GeneratePage: React.FC = () => {
|
||||
if (cancelled) return;
|
||||
if (status.status === "completed") {
|
||||
setCustomAudioUrl(status.output_audio_url);
|
||||
setCompletedTtsJobId(ttsJobId);
|
||||
setTtsJobId(null);
|
||||
setTtsError(null);
|
||||
message.success("语音合成完成!");
|
||||
@@ -351,6 +375,86 @@ const GeneratePage: React.FC = () => {
|
||||
});
|
||||
}, [customVoiceText, selectedVoice, synthesizeMutation]);
|
||||
|
||||
/* ── 存为素材 mutation ── */
|
||||
const saveToLibraryMutation = useMutation({
|
||||
mutationFn: (params: { name?: string; tag_ids?: string[] }) =>
|
||||
saveTtsToLibrary(completedTtsJobId!, params),
|
||||
onSuccess: () => {
|
||||
message.success({
|
||||
content: (
|
||||
<span>
|
||||
已保存到配音素材库!{" "}
|
||||
<a
|
||||
onClick={handleGoToLibrary}
|
||||
style={{
|
||||
color: "var(--primary-500, #6366f1)",
|
||||
cursor: "pointer",
|
||||
}}
|
||||
>
|
||||
去素材库查看
|
||||
</a>
|
||||
</span>
|
||||
),
|
||||
duration: 5,
|
||||
});
|
||||
setSaveModalOpen(false);
|
||||
setSaveName("");
|
||||
setSaveTagIds([]);
|
||||
setSaveNewTag("");
|
||||
setCompletedTtsJobId(null);
|
||||
setCustomAudioUrl(null);
|
||||
},
|
||||
onError: (err: Error) => {
|
||||
message.error(`保存失败:${err.message || "请重试"}`);
|
||||
},
|
||||
});
|
||||
|
||||
/** 打开存为素材弹窗 */
|
||||
const handleOpenSaveModal = useCallback(() => {
|
||||
setSaveName("");
|
||||
setSaveTagIds([]);
|
||||
setSaveNewTag("");
|
||||
setSaveModalOpen(true);
|
||||
}, []);
|
||||
|
||||
/** 确认保存 */
|
||||
const handleConfirmSave = useCallback(() => {
|
||||
if (!completedTtsJobId) return;
|
||||
saveToLibraryMutation.mutate({
|
||||
name: saveName.trim() || undefined,
|
||||
tag_ids: saveTagIds.length > 0 ? saveTagIds : undefined,
|
||||
});
|
||||
}, [completedTtsJobId, saveName, saveTagIds, saveToLibraryMutation]);
|
||||
|
||||
/** 在弹窗中新增标签(先创建再选中) */
|
||||
const handleAddTagInModal = useCallback(
|
||||
async (tagName: string) => {
|
||||
const trimmed = tagName.trim();
|
||||
if (!trimmed) return;
|
||||
/* 已在选中列表则跳过 */
|
||||
const existing = allTags.find((t) => t.name === trimmed);
|
||||
if (existing) {
|
||||
if (!saveTagIds.includes(existing.id)) {
|
||||
setSaveTagIds((prev) => [...prev, existing.id]);
|
||||
}
|
||||
return;
|
||||
}
|
||||
try {
|
||||
const created = await createTag(trimmed);
|
||||
setSaveTagIds((prev) => [...prev, created.id]);
|
||||
setSaveNewTag("");
|
||||
} catch {
|
||||
message.error(`创建标签"${trimmed}"失败`);
|
||||
}
|
||||
},
|
||||
[allTags, saveTagIds],
|
||||
);
|
||||
|
||||
/** 保存成功后跳转到素材库 */
|
||||
const handleGoToLibrary = useCallback(() => {
|
||||
navigate("/app/voice-materials");
|
||||
}, [navigate]);
|
||||
|
||||
const handleGenerate = useCallback(async () => {
|
||||
if (!title.trim()) {
|
||||
message.warning("请先选择或输入标题");
|
||||
@@ -910,16 +1014,136 @@ const GeneratePage: React.FC = () => {
|
||||
{ttsError}
|
||||
</Text>
|
||||
)}
|
||||
{customAudioUrl && (
|
||||
<Text
|
||||
{customAudioUrl && completedTtsJobId && (
|
||||
<div
|
||||
style={{
|
||||
color: "var(--success, #10b981)",
|
||||
marginTop: 8,
|
||||
display: "block",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
gap: 12,
|
||||
}}
|
||||
>
|
||||
✓ 语音合成完成
|
||||
</Text>
|
||||
<Text style={{ color: "var(--success, #10b981)" }}>
|
||||
✓ 语音合成完成
|
||||
</Text>
|
||||
<button
|
||||
className="xx-btn xx-btn-primary"
|
||||
style={{ height: 30, padding: "0 14px", fontSize: 12 }}
|
||||
onClick={handleOpenSaveModal}
|
||||
>
|
||||
<SaveOutlined /> 存为素材
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* ── 存为素材弹窗 ── */}
|
||||
{saveModalOpen && (
|
||||
<div
|
||||
className="xx-save-modal-overlay"
|
||||
onClick={() => setSaveModalOpen(false)}
|
||||
>
|
||||
<div
|
||||
className="xx-save-modal"
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
>
|
||||
<div className="xx-save-modal-header">
|
||||
<span>保存到配音素材库</span>
|
||||
<button
|
||||
className="xx-save-modal-close"
|
||||
onClick={() => setSaveModalOpen(false)}
|
||||
>
|
||||
<CloseOutlined />
|
||||
</button>
|
||||
</div>
|
||||
<div className="xx-save-modal-body">
|
||||
<label className="xx-save-modal-label">素材名称</label>
|
||||
<input
|
||||
className="xx-save-modal-input"
|
||||
placeholder="留空则自动生成名称"
|
||||
value={saveName}
|
||||
onChange={(e) => setSaveName(e.target.value)}
|
||||
maxLength={50}
|
||||
/>
|
||||
<label className="xx-save-modal-label">
|
||||
标签
|
||||
<span
|
||||
style={{
|
||||
fontWeight: 400,
|
||||
color: "var(--text-tertiary, #94a3b8)",
|
||||
}}
|
||||
>
|
||||
(可选)
|
||||
</span>
|
||||
</label>
|
||||
<div className="xx-save-modal-tags">
|
||||
{saveTagIds.map((id) => {
|
||||
const tag = allTags.find((t) => t.id === id);
|
||||
return tag ? (
|
||||
<span key={id} className="xx-save-modal-tag active">
|
||||
{tag.name}
|
||||
<CloseOutlined
|
||||
className="xx-save-modal-tag-remove"
|
||||
onClick={() =>
|
||||
setSaveTagIds((prev) =>
|
||||
prev.filter((x) => x !== id),
|
||||
)
|
||||
}
|
||||
/>
|
||||
</span>
|
||||
) : null;
|
||||
})}
|
||||
<input
|
||||
className="xx-save-modal-tag-input"
|
||||
placeholder="输入标签名回车添加"
|
||||
value={saveNewTag}
|
||||
onChange={(e) => setSaveNewTag(e.target.value)}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter") {
|
||||
e.preventDefault();
|
||||
handleAddTagInModal(saveNewTag);
|
||||
}
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
{allTags.length > 0 && (
|
||||
<div className="xx-save-modal-tag-presets">
|
||||
{allTags
|
||||
.filter((t) => !saveTagIds.includes(t.id))
|
||||
.slice(0, 12)
|
||||
.map((t) => (
|
||||
<button
|
||||
key={t.id}
|
||||
className="xx-save-modal-tag-preset"
|
||||
onClick={() =>
|
||||
setSaveTagIds((prev) => [...prev, t.id])
|
||||
}
|
||||
>
|
||||
{t.name}
|
||||
<PlusOutlined
|
||||
style={{ fontSize: 10, marginLeft: 4 }}
|
||||
/>
|
||||
</button>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
<div className="xx-save-modal-footer">
|
||||
<button
|
||||
className="xx-btn xx-btn-ghost"
|
||||
onClick={() => setSaveModalOpen(false)}
|
||||
>
|
||||
取消
|
||||
</button>
|
||||
<button
|
||||
className="xx-btn xx-btn-primary"
|
||||
disabled={saveToLibraryMutation.isPending}
|
||||
onClick={handleConfirmSave}
|
||||
>
|
||||
{saveToLibraryMutation.isPending ? "保存中…" : "保存"}
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
@@ -883,3 +883,194 @@
|
||||
flex-direction: column;
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 存为素材弹窗 ── */
|
||||
.xx-save-modal-overlay {
|
||||
position: fixed;
|
||||
inset: 0;
|
||||
background: rgba(0, 0, 0, 0.45);
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
z-index: 1000;
|
||||
animation: xxFadeIn 0.15s ease;
|
||||
}
|
||||
|
||||
.xx-save-modal {
|
||||
background: var(--bg-card, #fff);
|
||||
border-radius: var(--radius-lg, 16px);
|
||||
width: 420px;
|
||||
max-width: 90vw;
|
||||
box-shadow: 0 20px 60px rgba(0, 0, 0, 0.15);
|
||||
animation: xxSlideUp 0.2s ease;
|
||||
}
|
||||
|
||||
.xx-save-modal-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 16px 20px;
|
||||
border-bottom: 1px solid var(--border-light, #f1f5f9);
|
||||
font-weight: 600;
|
||||
font-size: 15px;
|
||||
color: var(--text-primary, #0f172a);
|
||||
}
|
||||
|
||||
.xx-save-modal-close {
|
||||
background: none;
|
||||
border: none;
|
||||
cursor: pointer;
|
||||
color: var(--text-tertiary, #94a3b8);
|
||||
font-size: 14px;
|
||||
padding: 4px;
|
||||
border-radius: var(--radius-sm, 6px);
|
||||
transition: all 0.15s;
|
||||
}
|
||||
|
||||
.xx-save-modal-close:hover {
|
||||
background: var(--bg-hover, #f8fafc);
|
||||
color: var(--text-primary, #0f172a);
|
||||
}
|
||||
|
||||
.xx-save-modal-body {
|
||||
padding: 20px;
|
||||
}
|
||||
|
||||
.xx-save-modal-label {
|
||||
display: block;
|
||||
font-size: 13px;
|
||||
font-weight: 500;
|
||||
color: var(--text-secondary, #475569);
|
||||
margin-bottom: 6px;
|
||||
margin-top: 14px;
|
||||
}
|
||||
|
||||
.xx-save-modal-label:first-child {
|
||||
margin-top: 0;
|
||||
}
|
||||
|
||||
.xx-save-modal-input {
|
||||
width: 100%;
|
||||
padding: 8px 12px;
|
||||
border: 1px solid var(--border-color, #e2e8f0);
|
||||
border-radius: var(--radius-sm, 10px);
|
||||
font-size: 14px;
|
||||
outline: none;
|
||||
background: var(--bg-surface, #fff);
|
||||
color: var(--text-primary, #0f172a);
|
||||
transition: border-color 0.15s;
|
||||
}
|
||||
|
||||
.xx-save-modal-input:focus {
|
||||
border-color: var(--primary-500, #6366f1);
|
||||
box-shadow: 0 0 0 2px var(--primary-100, #e0e7ff);
|
||||
}
|
||||
|
||||
.xx-save-modal-tags {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
gap: 6px;
|
||||
padding: 8px;
|
||||
border: 1px solid var(--border-color, #e2e8f0);
|
||||
border-radius: var(--radius-sm, 10px);
|
||||
min-height: 40px;
|
||||
align-items: center;
|
||||
cursor: text;
|
||||
transition: border-color 0.15s;
|
||||
}
|
||||
|
||||
.xx-save-modal-tags:focus-within {
|
||||
border-color: var(--primary-500, #6366f1);
|
||||
box-shadow: 0 0 0 2px var(--primary-100, #e0e7ff);
|
||||
}
|
||||
|
||||
.xx-save-modal-tag {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 4px;
|
||||
padding: 2px 8px;
|
||||
font-size: 12px;
|
||||
border-radius: 12px;
|
||||
background: var(--primary-50, #eef2ff);
|
||||
color: var(--primary-600, #4f46e5);
|
||||
border: 1px solid var(--primary-200, #c7d2fe);
|
||||
}
|
||||
|
||||
.xx-save-modal-tag-remove {
|
||||
font-size: 10px;
|
||||
cursor: pointer;
|
||||
opacity: 0.6;
|
||||
transition: opacity 0.15s;
|
||||
}
|
||||
|
||||
.xx-save-modal-tag-remove:hover {
|
||||
opacity: 1;
|
||||
}
|
||||
|
||||
.xx-save-modal-tag-input {
|
||||
border: none;
|
||||
outline: none;
|
||||
flex: 1;
|
||||
min-width: 100px;
|
||||
font-size: 13px;
|
||||
background: transparent;
|
||||
color: var(--text-primary, #0f172a);
|
||||
}
|
||||
|
||||
.xx-save-modal-tag-input::placeholder {
|
||||
color: var(--text-tertiary, #94a3b8);
|
||||
}
|
||||
|
||||
.xx-save-modal-tag-presets {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
gap: 6px;
|
||||
margin-top: 10px;
|
||||
padding-top: 10px;
|
||||
border-top: 1px solid var(--border-light, #f1f5f9);
|
||||
}
|
||||
|
||||
.xx-save-modal-tag-preset {
|
||||
padding: 3px 10px;
|
||||
font-size: 12px;
|
||||
border: 1px solid var(--border-color, #e2e8f0);
|
||||
border-radius: 12px;
|
||||
background: var(--bg-surface, #fff);
|
||||
color: var(--text-secondary, #475569);
|
||||
cursor: pointer;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
|
||||
.xx-save-modal-tag-preset:hover {
|
||||
border-color: var(--primary-300, #a5b4fc);
|
||||
color: var(--primary-600, #4f46e5);
|
||||
background: var(--primary-50, #eef2ff);
|
||||
}
|
||||
|
||||
.xx-save-modal-footer {
|
||||
display: flex;
|
||||
justify-content: flex-end;
|
||||
gap: 8px;
|
||||
padding: 14px 20px;
|
||||
border-top: 1px solid var(--border-light, #f1f5f9);
|
||||
}
|
||||
|
||||
@keyframes xxFadeIn {
|
||||
from {
|
||||
opacity: 0;
|
||||
}
|
||||
to {
|
||||
opacity: 1;
|
||||
}
|
||||
}
|
||||
|
||||
@keyframes xxSlideUp {
|
||||
from {
|
||||
transform: translateY(12px);
|
||||
opacity: 0;
|
||||
}
|
||||
to {
|
||||
transform: translateY(0);
|
||||
opacity: 1;
|
||||
}
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -111,7 +111,9 @@
|
||||
|
||||
.vmat-card.playing {
|
||||
border-color: var(--primary-500);
|
||||
box-shadow: 0 0 0 1px var(--primary-500), var(--shadow-md);
|
||||
box-shadow:
|
||||
0 0 0 1px var(--primary-500),
|
||||
var(--shadow-md);
|
||||
}
|
||||
|
||||
/* 性别色带 */
|
||||
@@ -812,3 +814,449 @@
|
||||
justify-content: space-between;
|
||||
}
|
||||
}
|
||||
|
||||
/* ─── 批量操作栏 ─────────────────────────────────────────── */
|
||||
|
||||
.vmat-batch-bar {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 8px 16px;
|
||||
margin-bottom: 16px;
|
||||
background: var(--primary-soft, #eef2ff);
|
||||
border: 1px solid var(--primary-color, #6366f1);
|
||||
border-radius: 8px;
|
||||
animation: vmat-batch-bar-in 0.2s ease;
|
||||
}
|
||||
|
||||
@keyframes vmat-batch-bar-in {
|
||||
from {
|
||||
opacity: 0;
|
||||
transform: translateY(-8px);
|
||||
}
|
||||
to {
|
||||
opacity: 1;
|
||||
transform: translateY(0);
|
||||
}
|
||||
}
|
||||
|
||||
.vmat-batch-bar-left {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 12px;
|
||||
}
|
||||
|
||||
.vmat-batch-bar-right {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
.vmat-select-all {
|
||||
font-size: 13px;
|
||||
color: var(--text-secondary, #64748b);
|
||||
cursor: pointer;
|
||||
user-select: none;
|
||||
}
|
||||
|
||||
.vmat-select-all:hover {
|
||||
color: var(--primary-color, #6366f1);
|
||||
}
|
||||
|
||||
.vmat-batch-count {
|
||||
font-size: 13px;
|
||||
font-weight: 500;
|
||||
color: var(--primary-color, #6366f1);
|
||||
}
|
||||
|
||||
/* ─── 自定义 checkbox ────────────────────────────────────── */
|
||||
|
||||
.vmat-checkbox {
|
||||
width: 18px;
|
||||
height: 18px;
|
||||
border: 2px solid var(--neutral-300, #cbd5e1);
|
||||
border-radius: 4px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
cursor: pointer;
|
||||
transition: all 0.15s ease;
|
||||
background: var(--bg-primary, #fff);
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
.vmat-checkbox:hover {
|
||||
border-color: var(--primary-color, #6366f1);
|
||||
}
|
||||
|
||||
.vmat-checkbox.checked {
|
||||
background: var(--primary-color, #6366f1);
|
||||
border-color: var(--primary-color, #6366f1);
|
||||
color: #fff;
|
||||
font-size: 10px;
|
||||
}
|
||||
|
||||
/* ─── 卡片 checkbox 覆盖 ─────────────────────────────────── */
|
||||
|
||||
.vmat-card-checkbox {
|
||||
position: absolute;
|
||||
top: 8px;
|
||||
left: 8px;
|
||||
z-index: 2;
|
||||
}
|
||||
|
||||
.vmat-row-checkbox {
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
/* ─── 选中态 + 批量模式 ──────────────────────────────────── */
|
||||
|
||||
.vmat-card.selected {
|
||||
border-color: var(--primary-color, #6366f1);
|
||||
box-shadow: 0 0 0 1px var(--primary-color, #6366f1);
|
||||
}
|
||||
|
||||
.vmat-row.selected {
|
||||
background: var(--primary-soft, #eef2ff);
|
||||
}
|
||||
|
||||
.vmat-card.batch-mode {
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.vmat-row.batch-mode {
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
/* ─── 进度条拖拽 thumb ───────────────────────────────────── */
|
||||
|
||||
.vmat-progress-thumb {
|
||||
position: absolute;
|
||||
top: 50%;
|
||||
transform: translate(-50%, -50%);
|
||||
width: 12px;
|
||||
height: 12px;
|
||||
border-radius: 50%;
|
||||
background: var(--primary-color, #6366f1);
|
||||
border: 2px solid #fff;
|
||||
box-shadow: 0 1px 3px rgba(0, 0, 0, 0.2);
|
||||
pointer-events: none;
|
||||
z-index: 1;
|
||||
}
|
||||
|
||||
/* ─── 音量控制 ────────────────────────────────────────────── */
|
||||
|
||||
.vmat-volume {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 4px;
|
||||
margin-top: 6px;
|
||||
}
|
||||
|
||||
.vmat-volume-btn {
|
||||
background: none;
|
||||
border: none;
|
||||
padding: 0;
|
||||
cursor: pointer;
|
||||
color: var(--text-secondary, #64748b);
|
||||
font-size: 14px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
transition: color 0.15s;
|
||||
}
|
||||
|
||||
.vmat-volume-btn:hover {
|
||||
color: var(--primary-color, #6366f1);
|
||||
}
|
||||
|
||||
.vmat-volume-slider {
|
||||
-webkit-appearance: none;
|
||||
appearance: none;
|
||||
width: 48px;
|
||||
height: 4px;
|
||||
border-radius: 2px;
|
||||
background: var(--neutral-200, #e2e8f0);
|
||||
outline: none;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.vmat-volume-slider::-webkit-slider-thumb {
|
||||
-webkit-appearance: none;
|
||||
appearance: none;
|
||||
width: 10px;
|
||||
height: 10px;
|
||||
border-radius: 50%;
|
||||
background: var(--primary-color, #6366f1);
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.vmat-volume-slider::-moz-range-thumb {
|
||||
width: 10px;
|
||||
height: 10px;
|
||||
border-radius: 50%;
|
||||
background: var(--primary-color, #6366f1);
|
||||
cursor: pointer;
|
||||
border: none;
|
||||
}
|
||||
|
||||
/* ─── 上传进度条 ──────────────────────────────────────────── */
|
||||
|
||||
.vmat-upload-progress {
|
||||
position: relative;
|
||||
height: 20px;
|
||||
background: var(--neutral-100, #f1f5f9);
|
||||
border-radius: 10px;
|
||||
overflow: hidden;
|
||||
margin-top: 12px;
|
||||
}
|
||||
|
||||
.vmat-upload-progress-bar {
|
||||
height: 100%;
|
||||
background: linear-gradient(90deg, var(--primary-color, #6366f1), #818cf8);
|
||||
border-radius: 10px;
|
||||
transition: width 0.3s ease;
|
||||
}
|
||||
|
||||
.vmat-upload-progress-text {
|
||||
position: absolute;
|
||||
top: 50%;
|
||||
left: 50%;
|
||||
transform: translate(-50%, -50%);
|
||||
font-size: 11px;
|
||||
font-weight: 600;
|
||||
color: var(--text-primary, #1e293b);
|
||||
}
|
||||
|
||||
/* ─── 批量打标签 Popover ──────────────────────────────────── */
|
||||
|
||||
.vmat-tag-popover {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
gap: 8px;
|
||||
max-width: 240px;
|
||||
}
|
||||
|
||||
.vmat-tag-pop-btn {
|
||||
padding: 4px 12px;
|
||||
border: 1px solid var(--neutral-200, #e2e8f0);
|
||||
border-radius: 14px;
|
||||
background: var(--bg-primary, #fff);
|
||||
font-size: 12px;
|
||||
color: var(--text-secondary, #64748b);
|
||||
cursor: pointer;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
|
||||
.vmat-tag-pop-btn:hover {
|
||||
border-color: var(--primary-color, #6366f1);
|
||||
color: var(--primary-color, #6366f1);
|
||||
background: var(--primary-soft, #eef2ff);
|
||||
}
|
||||
|
||||
/* ─── 列表 checkbox 列 ────────────────────────────────────── */
|
||||
|
||||
.vmat-lh-checkbox {
|
||||
width: 30px;
|
||||
}
|
||||
|
||||
.vmat-list.batch-mode .vmat-list-header,
|
||||
.vmat-list.batch-mode .vmat-row {
|
||||
grid-template-columns: 30px 48px 1fr 80px 160px 120px 60px 70px 80px;
|
||||
}
|
||||
|
||||
/* ─── 标签筛选药丸条 ─────────────────────────────────────── */
|
||||
|
||||
.vmat-tag-filter-bar {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
overflow-x: auto;
|
||||
padding: 10px 0;
|
||||
margin-bottom: 12px;
|
||||
scrollbar-width: thin;
|
||||
}
|
||||
|
||||
.vmat-tag-filter-bar::-webkit-scrollbar {
|
||||
height: 4px;
|
||||
}
|
||||
|
||||
.vmat-tag-filter-bar::-webkit-scrollbar-thumb {
|
||||
background: var(--neutral-300, #cbd5e1);
|
||||
border-radius: 2px;
|
||||
}
|
||||
|
||||
.vmat-filter-pill {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 4px;
|
||||
padding: 6px 16px;
|
||||
border: 1px solid var(--border-primary);
|
||||
border-radius: var(--radius-full, 9999px);
|
||||
background: var(--bg-surface);
|
||||
color: var(--text-secondary);
|
||||
font-size: 13px;
|
||||
cursor: pointer;
|
||||
white-space: nowrap;
|
||||
flex-shrink: 0;
|
||||
transition: all 0.2s ease;
|
||||
line-height: 1.4;
|
||||
}
|
||||
|
||||
.vmat-filter-pill:hover {
|
||||
border-color: var(--primary-300);
|
||||
color: var(--text-primary);
|
||||
background: var(--primary-50);
|
||||
}
|
||||
|
||||
.vmat-filter-pill:focus-visible {
|
||||
outline: 2px solid var(--primary-500);
|
||||
outline-offset: 2px;
|
||||
}
|
||||
|
||||
.vmat-filter-pill.active {
|
||||
background: var(--primary-500);
|
||||
color: #fff;
|
||||
border-color: var(--primary-500);
|
||||
}
|
||||
|
||||
.vmat-filter-pill.active .vmat-filter-pill-count {
|
||||
color: rgba(255, 255, 255, 0.75);
|
||||
}
|
||||
|
||||
.vmat-filter-pill-count {
|
||||
font-size: 11px;
|
||||
opacity: 0.7;
|
||||
color: var(--text-tertiary);
|
||||
}
|
||||
|
||||
/* ─── TagSelector ──────────────────────────────────────────── */
|
||||
|
||||
.vmat-tag-selector-wrapper {
|
||||
position: relative;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
.vmat-tag-selector {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
gap: 6px;
|
||||
padding: 8px;
|
||||
border: 1px solid var(--border-primary);
|
||||
border-radius: var(--radius-md);
|
||||
background: var(--bg-surface);
|
||||
min-height: 42px;
|
||||
cursor: text;
|
||||
transition:
|
||||
border-color 0.2s,
|
||||
box-shadow 0.2s;
|
||||
}
|
||||
|
||||
.vmat-tag-selector:focus-within {
|
||||
border-color: var(--primary-500);
|
||||
box-shadow: 0 0 0 2px var(--primary-100);
|
||||
}
|
||||
|
||||
.vmat-tag-selector-input {
|
||||
border: none;
|
||||
outline: none;
|
||||
flex: 1;
|
||||
min-width: 80px;
|
||||
font-size: 13px;
|
||||
background: transparent;
|
||||
color: var(--text-primary);
|
||||
}
|
||||
|
||||
.vmat-tag-selector-input::placeholder {
|
||||
color: var(--text-placeholder);
|
||||
}
|
||||
|
||||
.vmat-tag-suggestions {
|
||||
position: absolute;
|
||||
top: 100%;
|
||||
left: 0;
|
||||
right: 0;
|
||||
z-index: 10;
|
||||
margin-top: 4px;
|
||||
background: var(--bg-surface);
|
||||
border: 1px solid var(--border-primary);
|
||||
border-radius: var(--radius-md);
|
||||
box-shadow: var(--shadow-md, 0 4px 12px rgba(0, 0, 0, 0.1));
|
||||
max-height: 160px;
|
||||
overflow-y: auto;
|
||||
}
|
||||
|
||||
.vmat-tag-suggestion-item {
|
||||
padding: 8px 12px;
|
||||
font-size: 13px;
|
||||
color: var(--text-primary);
|
||||
cursor: pointer;
|
||||
transition: background 0.15s;
|
||||
}
|
||||
|
||||
.vmat-tag-suggestion-item:hover {
|
||||
background: var(--primary-50);
|
||||
}
|
||||
|
||||
.vmat-tag-selector-presets {
|
||||
width: 100%;
|
||||
display: flex;
|
||||
gap: 6px;
|
||||
flex-wrap: wrap;
|
||||
margin-top: 4px;
|
||||
padding-top: 8px;
|
||||
border-top: 1px solid var(--border-light);
|
||||
}
|
||||
|
||||
.vmat-tag-selector-preset {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 4px;
|
||||
padding: 2px 10px;
|
||||
font-size: 12px;
|
||||
border: 1px solid var(--border-primary);
|
||||
border-radius: 12px;
|
||||
background: var(--bg-surface);
|
||||
color: var(--text-secondary);
|
||||
cursor: pointer;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
|
||||
.vmat-tag-selector-preset:hover {
|
||||
border-color: var(--primary-300);
|
||||
color: var(--primary-600);
|
||||
background: var(--primary-50);
|
||||
}
|
||||
|
||||
.vmat-tag-selector-preset.selected {
|
||||
opacity: 0.4;
|
||||
pointer-events: none;
|
||||
}
|
||||
|
||||
/* ─── 标签溢出 + 空态 ─────────────────────────────────────── */
|
||||
|
||||
.vmat-tag-empty {
|
||||
font-size: 12px;
|
||||
color: var(--primary-500);
|
||||
cursor: pointer;
|
||||
padding: 2px 8px;
|
||||
border: 1px dashed var(--primary-300);
|
||||
border-radius: var(--radius-sm);
|
||||
transition: all 0.15s;
|
||||
}
|
||||
|
||||
.vmat-tag-empty:hover {
|
||||
background: var(--primary-50);
|
||||
border-color: var(--primary-500);
|
||||
}
|
||||
|
||||
.vmat-tag-overflow {
|
||||
font-style: italic;
|
||||
opacity: 0.7;
|
||||
cursor: default;
|
||||
}
|
||||
|
||||
/* ── 批量打标签 Popover 自定义输入 ── */
|
||||
.vmat-tag-pop-input-row {
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
|
||||
@@ -41,6 +41,10 @@ def __getattr__(name: str):
|
||||
from .tts_synthesis import process_tts_synthesis
|
||||
|
||||
return process_tts_synthesis
|
||||
elif name == "process_tts_segment_synthesis":
|
||||
from .tts_synthesis import process_tts_segment_synthesis
|
||||
|
||||
return process_tts_segment_synthesis
|
||||
elif name == "run_ai_recommend":
|
||||
from .ai_tasks import run_ai_recommend
|
||||
|
||||
@@ -62,6 +66,7 @@ __all__ = [
|
||||
"extract_background_task",
|
||||
"process_voice_clone",
|
||||
"process_tts_synthesis",
|
||||
"process_tts_segment_synthesis",
|
||||
"run_ai_recommend",
|
||||
"run_generate_cover",
|
||||
]
|
||||
|
||||
@@ -104,3 +104,77 @@ def process_tts_synthesis(self: Task, job_id: str) -> dict:
|
||||
finally:
|
||||
if session is not None:
|
||||
session.close()
|
||||
|
||||
|
||||
@celery_app.task(bind=True, max_retries=2, name="worker.process_tts_segment_synthesis")
|
||||
def process_tts_segment_synthesis(self: Task, job_id: str) -> dict:
|
||||
"""分段合成轮询任务 — 轮询多个 CosyVoice 子任务并合并音频。
|
||||
|
||||
与 process_tts_synthesis 类似,但超时更长(300s),
|
||||
因为分段任务需要等待所有子任务完成。
|
||||
"""
|
||||
session = None
|
||||
try:
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyTTSJobRepository(session)
|
||||
workflow = TTSWorkflowService(
|
||||
repository=repo,
|
||||
cosyvoice_service=CosyVoiceService(),
|
||||
)
|
||||
|
||||
updated_job = workflow.poll_and_process_synthesis(job_id, timeout=300)
|
||||
session.commit()
|
||||
|
||||
logger.info(f"TTS segment synthesis completed: job_id={job_id}, " f"audio_url={updated_job.output_audio_url}")
|
||||
return {
|
||||
"ok": True,
|
||||
"job_id": job_id,
|
||||
"audio_url": updated_job.output_audio_url,
|
||||
}
|
||||
|
||||
except Retry:
|
||||
raise
|
||||
|
||||
except CosyVoiceTimeoutError as e:
|
||||
logger.warning(f"TTS segment synthesis timeout for {job_id}: {e}")
|
||||
if session is not None:
|
||||
session.rollback()
|
||||
raise self.retry(exc=e, countdown=60)
|
||||
|
||||
except CosyVoiceError as e:
|
||||
logger.error(f"TTS segment synthesis failed for {job_id}: {e}")
|
||||
if session is not None:
|
||||
session.rollback()
|
||||
try:
|
||||
if session is not None:
|
||||
job = repo.get(job_id)
|
||||
if job is not None:
|
||||
job.mark_failed(str(e))
|
||||
repo.update(job)
|
||||
session.commit()
|
||||
except Exception as inner_e:
|
||||
logger.error(f"Failed to mark job as failed: {inner_e}")
|
||||
if session is not None:
|
||||
session.rollback()
|
||||
return {"ok": False, "job_id": job_id, "error": str(e)}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"TTS segment synthesis unexpected error for {job_id}: {e}")
|
||||
if session is not None:
|
||||
session.rollback()
|
||||
try:
|
||||
if session is not None:
|
||||
job = repo.get(job_id)
|
||||
if job is not None:
|
||||
job.mark_failed(str(e))
|
||||
repo.update(job)
|
||||
session.commit()
|
||||
except Exception as inner_e:
|
||||
logger.error(f"Failed to mark job as failed: {inner_e}")
|
||||
if session is not None:
|
||||
session.rollback()
|
||||
return {"ok": False, "job_id": job_id, "error": str(e)}
|
||||
|
||||
finally:
|
||||
if session is not None:
|
||||
session.close()
|
||||
|
||||
@@ -95,6 +95,39 @@
|
||||
"id"
|
||||
]
|
||||
},
|
||||
"asset_tags": {
|
||||
"columns": [
|
||||
{
|
||||
"index": false,
|
||||
"name": "asset_id",
|
||||
"nullable": false,
|
||||
"primary_key": true,
|
||||
"type": "VARCHAR(36)",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"index": false,
|
||||
"name": "tag_id",
|
||||
"nullable": false,
|
||||
"primary_key": true,
|
||||
"type": "VARCHAR(36)",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"index": false,
|
||||
"name": "created_at",
|
||||
"nullable": false,
|
||||
"primary_key": false,
|
||||
"type": "DATETIME",
|
||||
"unique": false
|
||||
}
|
||||
],
|
||||
"indexes": [],
|
||||
"primary_key": [
|
||||
"asset_id",
|
||||
"tag_id"
|
||||
]
|
||||
},
|
||||
"assets": {
|
||||
"columns": [
|
||||
{
|
||||
@@ -2056,6 +2089,54 @@
|
||||
"id"
|
||||
]
|
||||
},
|
||||
"tags": {
|
||||
"columns": [
|
||||
{
|
||||
"index": false,
|
||||
"name": "id",
|
||||
"nullable": false,
|
||||
"primary_key": true,
|
||||
"type": "VARCHAR(36)",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"index": true,
|
||||
"name": "user_id",
|
||||
"nullable": false,
|
||||
"primary_key": false,
|
||||
"type": "VARCHAR(36)",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"index": false,
|
||||
"name": "name",
|
||||
"nullable": false,
|
||||
"primary_key": false,
|
||||
"type": "VARCHAR(100)",
|
||||
"unique": false
|
||||
},
|
||||
{
|
||||
"index": false,
|
||||
"name": "created_at",
|
||||
"nullable": false,
|
||||
"primary_key": false,
|
||||
"type": "DATETIME",
|
||||
"unique": false
|
||||
}
|
||||
],
|
||||
"indexes": [
|
||||
{
|
||||
"columns": [
|
||||
"user_id"
|
||||
],
|
||||
"name": "ix_tags_user_id",
|
||||
"unique": false
|
||||
}
|
||||
],
|
||||
"primary_key": [
|
||||
"id"
|
||||
]
|
||||
},
|
||||
"template_categories": {
|
||||
"columns": [
|
||||
{
|
||||
|
||||
@@ -42,3 +42,37 @@ class InMemoryAssetRepository:
|
||||
del self._assets[asset_id]
|
||||
return True
|
||||
return False
|
||||
|
||||
def batch_delete(self, asset_ids: list[str]) -> int:
|
||||
"""批量删除素材,返回实际删除数量。"""
|
||||
count = 0
|
||||
for aid in asset_ids:
|
||||
if aid in self._assets:
|
||||
del self._assets[aid]
|
||||
count += 1
|
||||
return count
|
||||
|
||||
def find_by_project(
|
||||
self,
|
||||
project_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
) -> list[Asset]:
|
||||
items = [a for a in self._assets.values() if a.project_id == project_id]
|
||||
return items[skip : skip + limit]
|
||||
|
||||
def find_by_id(self, asset_id: str) -> Asset | None:
|
||||
return self._assets.get(asset_id)
|
||||
|
||||
def find_by_tag_ids(
|
||||
self,
|
||||
tag_ids: list[str],
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
) -> list[Asset]:
|
||||
"""查找包含所有指定标签的素材。"""
|
||||
if not tag_ids:
|
||||
return []
|
||||
tag_set = set(tag_ids)
|
||||
items = [a for a in self._assets.values() if tag_set.issubset(set(a.tag_ids))]
|
||||
return items[skip : skip + limit]
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
"""标签 InMemory 仓储实现。"""
|
||||
|
||||
from packages.domain import Tag
|
||||
|
||||
|
||||
class InMemoryTagRepository:
|
||||
def __init__(self):
|
||||
self._tags: dict[str, Tag] = {}
|
||||
|
||||
def create(self, tag: Tag) -> Tag:
|
||||
self._tags[tag.id] = tag
|
||||
return tag
|
||||
|
||||
def get(self, tag_id: str) -> Tag | None:
|
||||
return self._tags.get(tag_id)
|
||||
|
||||
def find_by_name(self, user_id: str, name: str) -> Tag | None:
|
||||
for tag in self._tags.values():
|
||||
if tag.user_id == user_id and tag.name == name:
|
||||
return tag
|
||||
return None
|
||||
|
||||
def list_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
) -> list[Tag]:
|
||||
tags = [tag for tag in self._tags.values() if tag.user_id == user_id]
|
||||
tags.sort(key=lambda t: t.created_at, reverse=True)
|
||||
return tags[skip : skip + limit]
|
||||
|
||||
def count_by_user(self, user_id: str) -> int:
|
||||
return sum(1 for tag in self._tags.values() if tag.user_id == user_id)
|
||||
|
||||
def delete(self, tag_id: str) -> bool:
|
||||
if tag_id in self._tags:
|
||||
del self._tags[tag_id]
|
||||
return True
|
||||
return False
|
||||
@@ -3,7 +3,7 @@ from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetModel
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetModel, AssetTagModel
|
||||
from packages.domain import Asset, AssetStatus, ClassificationStatus
|
||||
|
||||
|
||||
@@ -87,6 +87,8 @@ class SQLAlchemyAssetRepository:
|
||||
updated_at=now,
|
||||
)
|
||||
self.session.add(model)
|
||||
self.session.flush()
|
||||
self._sync_asset_tags(asset.id, asset.tag_ids)
|
||||
self.session.commit()
|
||||
return asset
|
||||
|
||||
@@ -109,6 +111,8 @@ class SQLAlchemyAssetRepository:
|
||||
model.quality_score = asset.quality_score
|
||||
model.uploaded_by_user_id = asset.uploaded_by_user_id or model.uploaded_by_user_id
|
||||
model.updated_at = datetime.now(timezone.utc)
|
||||
self.session.flush()
|
||||
self._sync_asset_tags(asset.id, asset.tag_ids)
|
||||
self.session.commit()
|
||||
return asset
|
||||
|
||||
@@ -120,6 +124,14 @@ class SQLAlchemyAssetRepository:
|
||||
return True
|
||||
return False
|
||||
|
||||
def batch_delete(self, asset_ids: list[str]) -> int:
|
||||
"""批量删除素材,返回实际删除数量。"""
|
||||
if not asset_ids:
|
||||
return 0
|
||||
count = self.session.query(AssetModel).filter(AssetModel.id.in_(asset_ids)).delete(synchronize_session=False)
|
||||
self.session.commit()
|
||||
return count
|
||||
|
||||
def count_by_project(self, project_id: str) -> int:
|
||||
return self.session.query(AssetModel).filter(AssetModel.project_id == project_id).count()
|
||||
|
||||
@@ -195,6 +207,11 @@ class SQLAlchemyAssetRepository:
|
||||
"audio": "audio/mpeg",
|
||||
"image": "image/jpeg",
|
||||
}.get(mime_type, mime_type)
|
||||
# 查询关联的 tag_ids
|
||||
tag_ids = [
|
||||
row.tag_id
|
||||
for row in self.session.query(AssetTagModel.tag_id).filter(AssetTagModel.asset_id == model.id).all()
|
||||
]
|
||||
return Asset(
|
||||
id=model.id,
|
||||
project_id=model.project_id,
|
||||
@@ -214,6 +231,39 @@ class SQLAlchemyAssetRepository:
|
||||
quality_score=model.quality_score,
|
||||
uploaded_by_user_id=model.uploaded_by_user_id,
|
||||
metadata=metadata,
|
||||
tag_ids=tag_ids,
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
def _sync_asset_tags(self, asset_id: str, tag_ids: list[str]) -> None:
|
||||
"""同步素材-标签关联表(全量替换)。"""
|
||||
self.session.query(AssetTagModel).filter(AssetTagModel.asset_id == asset_id).delete(synchronize_session=False)
|
||||
for tag_id in tag_ids:
|
||||
self.session.add(AssetTagModel(asset_id=asset_id, tag_id=tag_id))
|
||||
|
||||
def find_by_tag_ids(
|
||||
self,
|
||||
tag_ids: list[str],
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
) -> list[Asset]:
|
||||
"""查找包含所有指定标签的素材。"""
|
||||
if not tag_ids:
|
||||
return []
|
||||
from sqlalchemy import func
|
||||
|
||||
# 找出同时拥有所有指定 tag_id 的 asset_id
|
||||
tag_set = set(tag_ids)
|
||||
asset_ids = (
|
||||
self.session.query(AssetTagModel.asset_id)
|
||||
.filter(AssetTagModel.tag_id.in_(tag_set))
|
||||
.group_by(AssetTagModel.asset_id)
|
||||
.having(func.count(AssetTagModel.tag_id) == len(tag_set))
|
||||
.all()
|
||||
)
|
||||
ids = [row[0] for row in asset_ids]
|
||||
if not ids:
|
||||
return []
|
||||
models = self.session.query(AssetModel).filter(AssetModel.id.in_(ids)).offset(skip).limit(limit).all()
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
@@ -90,6 +90,29 @@ class AssetModel(Base):
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class TagModel(Base):
|
||||
"""标签 ORM 模型。"""
|
||||
|
||||
__tablename__ = "tags"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
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))
|
||||
|
||||
__table_args__ = (UniqueConstraint("user_id", "name", name="uq_tags_user_name"),)
|
||||
|
||||
|
||||
class AssetTagModel(Base):
|
||||
"""素材-标签关联表 ORM 模型。"""
|
||||
|
||||
__tablename__ = "asset_tags"
|
||||
|
||||
asset_id = Column(String(36), primary_key=True)
|
||||
tag_id = Column(String(36), primary_key=True)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class EditTemplateModel(Base):
|
||||
"""Phase 8 剪辑模板 ORM 模型
|
||||
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
"""标签 SQLAlchemy 仓储实现。"""
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetTagModel, TagModel
|
||||
from packages.domain import Tag
|
||||
|
||||
|
||||
class SQLAlchemyTagRepository:
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def create(self, tag: Tag) -> Tag:
|
||||
model = TagModel(
|
||||
id=tag.id,
|
||||
user_id=tag.user_id,
|
||||
name=tag.name,
|
||||
created_at=tag.created_at,
|
||||
)
|
||||
self.session.add(model)
|
||||
self.session.commit()
|
||||
return tag
|
||||
|
||||
def get(self, tag_id: str) -> Tag | None:
|
||||
model = self.session.query(TagModel).filter(TagModel.id == tag_id).first()
|
||||
if model is None:
|
||||
return None
|
||||
return self._to_domain(model)
|
||||
|
||||
def find_by_name(self, user_id: str, name: str) -> Tag | None:
|
||||
model = self.session.query(TagModel).filter(TagModel.user_id == user_id, TagModel.name == name).first()
|
||||
if model is None:
|
||||
return None
|
||||
return self._to_domain(model)
|
||||
|
||||
def list_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
) -> list[Tag]:
|
||||
models = (
|
||||
self.session.query(TagModel)
|
||||
.filter(TagModel.user_id == user_id)
|
||||
.order_by(TagModel.created_at.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def count_by_user(self, user_id: str) -> int:
|
||||
return self.session.query(TagModel).filter(TagModel.user_id == user_id).count()
|
||||
|
||||
def delete(self, tag_id: str) -> bool:
|
||||
# 先清理关联表
|
||||
self.session.query(AssetTagModel).filter(AssetTagModel.tag_id == tag_id).delete(synchronize_session=False)
|
||||
model = self.session.query(TagModel).filter(TagModel.id == tag_id).first()
|
||||
if model is None:
|
||||
self.session.commit()
|
||||
return False
|
||||
self.session.delete(model)
|
||||
self.session.commit()
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _to_domain(model: TagModel) -> Tag:
|
||||
return Tag(
|
||||
id=model.id,
|
||||
user_id=model.user_id,
|
||||
name=model.name,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
@@ -0,0 +1,96 @@
|
||||
"""FFmpeg 音频合并器 — P1 长文本分段合成。
|
||||
|
||||
将多个分段音频文件合并为一个完整音频文件。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import tempfile
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AudioMergeError(Exception):
|
||||
"""音频合并异常。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class AudioMerger:
|
||||
"""使用 FFmpeg 合并多个音频文件。"""
|
||||
|
||||
def merge(self, audio_paths: list[str], output_format: str = "mp3") -> bytes:
|
||||
"""合并多个音频文件,返回合并后的音频数据。
|
||||
|
||||
使用 FFmpeg concat demuxer 按顺序拼接音频。
|
||||
所有输入文件必须为相同格式和采样率。
|
||||
|
||||
Args:
|
||||
audio_paths: 音频文件路径列表(按合成顺序)
|
||||
output_format: 输出格式(mp3/wav/pcm)
|
||||
|
||||
Returns:
|
||||
合并后的音频文件字节数据
|
||||
|
||||
Raises:
|
||||
AudioMergeError: 合并失败
|
||||
"""
|
||||
if not audio_paths:
|
||||
raise AudioMergeError("没有可合并的音频文件")
|
||||
|
||||
if len(audio_paths) == 1:
|
||||
with open(audio_paths[0], "rb") as f:
|
||||
return f.read()
|
||||
|
||||
temp_dir = tempfile.mkdtemp(prefix="tts_merge_")
|
||||
try:
|
||||
# 生成 concat demuxer 列表文件
|
||||
list_path = os.path.join(temp_dir, "concat_list.txt")
|
||||
with open(list_path, "w") as f:
|
||||
for path in audio_paths:
|
||||
# FFmpeg concat 文件需要 file: 前缀,路径中的 ' 和 \n 需转义
|
||||
escaped = path.replace("'", "'\\''").replace("\n", "\\n")
|
||||
f.write(f"file '{escaped}'\n")
|
||||
|
||||
output_path = os.path.join(temp_dir, f"merged.{output_format}")
|
||||
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"concat",
|
||||
"-safe",
|
||||
"0",
|
||||
"-i",
|
||||
list_path,
|
||||
"-c",
|
||||
"copy",
|
||||
output_path,
|
||||
]
|
||||
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=120,
|
||||
)
|
||||
|
||||
if result.returncode != 0:
|
||||
logger.error(f"FFmpeg 合并失败: stderr={result.stderr}")
|
||||
raise AudioMergeError(f"FFmpeg 合并失败: {result.stderr[:500]}")
|
||||
|
||||
with open(output_path, "rb") as f:
|
||||
return f.read()
|
||||
|
||||
except subprocess.TimeoutExpired:
|
||||
raise AudioMergeError("FFmpeg 合并超时(120 秒)")
|
||||
except AudioMergeError:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise AudioMergeError(f"音频合并失败: {e}")
|
||||
finally:
|
||||
shutil.rmtree(temp_dir, ignore_errors=True)
|
||||
@@ -0,0 +1,251 @@
|
||||
"""P2: TTS 流式合成服务 — WebSocket 实时音频推送。
|
||||
|
||||
通过 WebSocket 将合成音频以二进制帧实时推送给客户端。
|
||||
- 短文本(≤500 字):合成完整音频后分块推送
|
||||
- 长文本(>500 字):分段并发合成,逐段推送音频
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Any, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
|
||||
from packages.application.tts_job.text_splitter import split_text
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# WebSocket 二进制帧块大小(4KB)
|
||||
_AUDIO_CHUNK_SIZE = 4096
|
||||
# 分段并发上限
|
||||
_MAX_STREAMING_SEGMENT_WORKERS = 5
|
||||
# 长文本分段阈值
|
||||
_SEGMENT_THRESHOLD = 500
|
||||
# WebSocket 最大文本长度
|
||||
_MAX_TEXT_LENGTH = 10000
|
||||
|
||||
|
||||
class TTSStreamingError(Exception):
|
||||
"""TTS 流式合成异常。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class TTSStreamingService:
|
||||
"""TTS 流式合成服务。
|
||||
|
||||
通过 WebSocket 实时推送合成音频。
|
||||
使用 CosyVoiceService(同步 REST API)合成,
|
||||
通过 asyncio.to_thread 桥接到异步 WebSocket。
|
||||
"""
|
||||
|
||||
def __init__(self, cosyvoice_service: CosyVoiceService) -> None:
|
||||
self._cosyvoice = cosyvoice_service
|
||||
|
||||
async def synthesize_and_stream(self, websocket: Any, params: dict) -> None:
|
||||
"""根据文本长度选择流式合成策略。
|
||||
|
||||
Args:
|
||||
websocket: FastAPI WebSocket 连接
|
||||
params: 合成参数(text, voice_id, sample_rate, format, speed)
|
||||
"""
|
||||
text = params.get("text", "")
|
||||
if not text:
|
||||
await self._send_json(websocket, {"type": "error", "message": "文本不能为空"})
|
||||
return
|
||||
|
||||
if len(text) > _MAX_TEXT_LENGTH:
|
||||
await self._send_json(
|
||||
websocket,
|
||||
{"type": "error", "message": f"文本过长,最大 {_MAX_TEXT_LENGTH} 字"},
|
||||
)
|
||||
return
|
||||
|
||||
if len(text) <= _SEGMENT_THRESHOLD:
|
||||
await self._stream_short_text(websocket, params)
|
||||
else:
|
||||
await self._stream_long_text(websocket, params)
|
||||
|
||||
# ── 短文本流式合成 ────────────────────────────────────────
|
||||
|
||||
async def _stream_short_text(self, websocket: Any, params: dict) -> None:
|
||||
"""短文本:合成完整音频后分块推送。"""
|
||||
text = params["text"]
|
||||
voice_id = params.get("voice_id", "")
|
||||
sample_rate = params.get("sample_rate", 0)
|
||||
audio_format = params.get("format", "mp3")
|
||||
speed = params.get("speed", 1.0)
|
||||
|
||||
await self._send_json(
|
||||
websocket,
|
||||
{"type": "started", "segment_count": 1, "total_segments": 1},
|
||||
)
|
||||
|
||||
# 在线程池中执行同步合成
|
||||
try:
|
||||
result = await asyncio.to_thread(
|
||||
self._cosyvoice.submit_synthesize_task,
|
||||
text=text,
|
||||
voice_id=voice_id,
|
||||
sample_rate=sample_rate,
|
||||
format=audio_format,
|
||||
speed=speed,
|
||||
)
|
||||
except CosyVoiceError as e:
|
||||
logger.error(f"流式合成失败: {e}")
|
||||
await self._send_json(websocket, {"type": "error", "message": str(e)})
|
||||
return
|
||||
except Exception as e:
|
||||
logger.error(f"流式合成意外错误: {e}")
|
||||
await self._send_json(websocket, {"type": "error", "message": f"合成失败: {e}"})
|
||||
return
|
||||
|
||||
audio_url = result.get("audio_url", "")
|
||||
if not audio_url:
|
||||
await self._send_json(websocket, {"type": "error", "message": "合成未返回音频 URL"})
|
||||
return
|
||||
|
||||
# 下载并流式推送音频
|
||||
try:
|
||||
audio_data = await asyncio.to_thread(self._download_audio, audio_url)
|
||||
total_bytes = await self._stream_audio_chunks(websocket, audio_data)
|
||||
|
||||
await self._send_json(
|
||||
websocket,
|
||||
{
|
||||
"type": "done",
|
||||
"duration": result.get("duration", 0.0),
|
||||
"file_size": total_bytes,
|
||||
"format": audio_format,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"音频流式推送失败: {e}")
|
||||
await self._send_json(websocket, {"type": "error", "message": f"音频推送失败: {e}"})
|
||||
|
||||
# ── 长文本分段流式合成 ────────────────────────────────────
|
||||
|
||||
async def _stream_long_text(self, websocket: Any, params: dict) -> None:
|
||||
"""长文本:分段并发合成,逐段推送音频。"""
|
||||
text = params["text"]
|
||||
voice_id = params.get("voice_id", "")
|
||||
sample_rate = params.get("sample_rate", 0)
|
||||
audio_format = params.get("format", "mp3")
|
||||
speed = params.get("speed", 1.0)
|
||||
|
||||
segments = split_text(text, max_chars=_SEGMENT_THRESHOLD)
|
||||
segment_count = len(segments)
|
||||
|
||||
logger.info(f"流式分段合成: 原文={len(text)}字, 段数={segment_count}")
|
||||
|
||||
await self._send_json(
|
||||
websocket,
|
||||
{"type": "started", "segment_count": segment_count, "total_segments": segment_count},
|
||||
)
|
||||
|
||||
# 并发合成所有分段,按顺序流式推送
|
||||
queue: asyncio.Queue[tuple[int, Optional[bytes], Optional[str]]] = asyncio.Queue()
|
||||
completed_count = 0
|
||||
|
||||
async def _synthesize_one(idx: int, seg_text: str) -> None:
|
||||
"""合成单个分段并放入队列。"""
|
||||
try:
|
||||
result = await asyncio.to_thread(
|
||||
self._cosyvoice.submit_synthesize_task,
|
||||
text=seg_text,
|
||||
voice_id=voice_id,
|
||||
sample_rate=sample_rate,
|
||||
format=audio_format,
|
||||
speed=speed,
|
||||
)
|
||||
audio_url = result.get("audio_url", "")
|
||||
if audio_url:
|
||||
audio_data = await asyncio.to_thread(self._download_audio, audio_url)
|
||||
await queue.put((idx, audio_data, None))
|
||||
else:
|
||||
await queue.put((idx, None, "合成未返回音频 URL"))
|
||||
except Exception as e:
|
||||
await queue.put((idx, None, str(e)))
|
||||
|
||||
# 启动并发合成任务
|
||||
workers = [asyncio.create_task(_synthesize_one(idx, seg)) for idx, seg in enumerate(segments)]
|
||||
|
||||
# 按顺序消费队列,流式推送
|
||||
total_bytes = 0
|
||||
total_duration = 0.0
|
||||
consumed = 0
|
||||
|
||||
try:
|
||||
while consumed < segment_count:
|
||||
idx, audio_data, error = await queue.get()
|
||||
consumed += 1
|
||||
|
||||
if error:
|
||||
logger.error(f"分段 {idx + 1} 合成失败: {error}")
|
||||
await self._send_json(
|
||||
websocket,
|
||||
{"type": "error", "message": f"分段 {idx + 1} 合成失败: {error}"},
|
||||
)
|
||||
# 取消剩余 worker
|
||||
for w in workers:
|
||||
w.cancel()
|
||||
return
|
||||
|
||||
if audio_data:
|
||||
seg_bytes = await self._stream_audio_chunks(websocket, audio_data)
|
||||
total_bytes += seg_bytes
|
||||
|
||||
await self._send_json(
|
||||
websocket,
|
||||
{"type": "segment_done", "segment": idx + 1, "total": segment_count},
|
||||
)
|
||||
|
||||
# 等待所有 worker 完成
|
||||
await asyncio.gather(*workers, return_exceptions=True)
|
||||
|
||||
await self._send_json(
|
||||
websocket,
|
||||
{
|
||||
"type": "done",
|
||||
"duration": total_duration,
|
||||
"file_size": total_bytes,
|
||||
"format": audio_format,
|
||||
},
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"流式分段推送失败: {e}")
|
||||
await self._send_json(websocket, {"type": "error", "message": f"推送失败: {e}"})
|
||||
for w in workers:
|
||||
w.cancel()
|
||||
|
||||
# ── 工具方法 ────────────────────────────────────────────
|
||||
|
||||
def _download_audio(self, url: str) -> bytes:
|
||||
"""下载音频数据。"""
|
||||
resp = httpx.get(url, timeout=60.0, follow_redirects=True)
|
||||
resp.raise_for_status()
|
||||
return resp.content
|
||||
|
||||
async def _stream_audio_chunks(self, websocket: Any, audio_data: bytes) -> int:
|
||||
"""将音频数据分块通过 WebSocket 推送。
|
||||
|
||||
Returns:
|
||||
推送的总字节数
|
||||
"""
|
||||
total = 0
|
||||
for offset in range(0, len(audio_data), _AUDIO_CHUNK_SIZE):
|
||||
chunk = audio_data[offset : offset + _AUDIO_CHUNK_SIZE]
|
||||
await websocket.send_bytes(chunk)
|
||||
total += len(chunk)
|
||||
return total
|
||||
|
||||
async def _send_json(self, websocket: Any, data: dict) -> None:
|
||||
"""安全发送 JSON 帧。"""
|
||||
try:
|
||||
await websocket.send_json(data)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -0,0 +1,70 @@
|
||||
"""长文本分段工具 — P1 长文本分段合成。
|
||||
|
||||
将超过阈值的文本按句子边界分段,供 CosyVoice 并发合成后合并。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
# 中文句子结束符(含全角/半角)
|
||||
_SENTENCE_ENDS = frozenset("。!?;\n.!?;")
|
||||
|
||||
|
||||
def split_text(text: str, max_chars: int = 500) -> list[str]:
|
||||
"""将文本分段,每段不超过 max_chars 个字符。
|
||||
|
||||
优先在句子边界(句号、问号、感叹号、换行符)处分段。
|
||||
若单个句子超过 max_chars,则在逗号等次级标点处拆分。
|
||||
若仍超长,则硬切。
|
||||
|
||||
Args:
|
||||
text: 待分段文本
|
||||
max_chars: 每段最大字符数
|
||||
|
||||
Returns:
|
||||
分段列表,每段 ≤ max_chars。文本为空时返回空列表。
|
||||
"""
|
||||
text = text.strip()
|
||||
if not text:
|
||||
return []
|
||||
if len(text) <= max_chars:
|
||||
return [text]
|
||||
|
||||
segments: list[str] = []
|
||||
current = ""
|
||||
|
||||
for char in text:
|
||||
current += char
|
||||
if char in _SENTENCE_ENDS and len(current) >= 50:
|
||||
# 句子边界且长度合理,切段
|
||||
segments.append(current.strip())
|
||||
current = ""
|
||||
elif len(current) >= max_chars:
|
||||
# 达到上限,强制切段
|
||||
segments.append(current.strip())
|
||||
current = ""
|
||||
|
||||
if current.strip():
|
||||
segments.append(current.strip())
|
||||
|
||||
# 合并过短的段(< 50 字符且不是最后一段),减少 API 调用次数
|
||||
merged: list[str] = []
|
||||
buffer = ""
|
||||
for seg in segments:
|
||||
if buffer:
|
||||
combined = buffer + seg
|
||||
if len(combined) <= max_chars:
|
||||
buffer = combined
|
||||
continue
|
||||
merged.append(buffer)
|
||||
buffer = ""
|
||||
if len(seg) < 50:
|
||||
buffer = seg
|
||||
else:
|
||||
merged.append(seg)
|
||||
if buffer:
|
||||
if merged and len(merged[-1]) + len(buffer) <= max_chars:
|
||||
merged[-1] = merged[-1] + buffer
|
||||
else:
|
||||
merged.append(buffer)
|
||||
|
||||
return [s for s in merged if s]
|
||||
@@ -9,19 +9,35 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from packages.application.cosyvoice_service import (
|
||||
CosyVoiceAuthError,
|
||||
CosyVoiceError,
|
||||
CosyVoiceService,
|
||||
)
|
||||
from packages.application.tts_job.audio_merger import AudioMergeError, AudioMerger
|
||||
from packages.application.tts_job.text_splitter import split_text
|
||||
from packages.domain.tts_job import TTSJob, TTSJobStatus
|
||||
from packages.ports.tts_job_repository import TTSJobRepository
|
||||
from packages.shared.storage import SharedStorageService, get_shared_storage_service
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 长文本分段阈值:超过此字符数自动分段合成
|
||||
_SEGMENT_THRESHOLD = 500
|
||||
# 分段并发上限
|
||||
_MAX_SEGMENT_WORKERS = 5
|
||||
|
||||
|
||||
class TTSWorkflowError(Exception):
|
||||
"""TTS 合成工作流异常。"""
|
||||
@@ -46,9 +62,55 @@ class TTSWorkflowService:
|
||||
self,
|
||||
repository: TTSJobRepository,
|
||||
cosyvoice_service: CosyVoiceService,
|
||||
storage_service: Optional[SharedStorageService] = None,
|
||||
) -> None:
|
||||
self.repository = repository
|
||||
self.cosyvoice_service = cosyvoice_service
|
||||
self._storage_service = storage_service
|
||||
|
||||
@property
|
||||
def _storage(self) -> SharedStorageService:
|
||||
if self._storage_service is None:
|
||||
self._storage_service = get_shared_storage_service()
|
||||
return self._storage_service
|
||||
|
||||
def _transfer_audio_to_oss(
|
||||
self,
|
||||
temp_url: str,
|
||||
user_id: str,
|
||||
job_id: str,
|
||||
audio_format: str = "mp3",
|
||||
) -> tuple[str, str]:
|
||||
"""下载 CosyVoice 临时音频并转存到 OSS。
|
||||
|
||||
Returns:
|
||||
(permanent_url, storage_key) 元组。
|
||||
转存失败时回退到原始临时 URL,storage_key 为空字符串。
|
||||
"""
|
||||
storage_key = f"tts-outputs/{user_id}/{job_id}.{audio_format}"
|
||||
content_type_map = {
|
||||
"mp3": "audio/mpeg",
|
||||
"wav": "audio/wav",
|
||||
"pcm": "audio/pcm",
|
||||
"opus": "audio/opus",
|
||||
}
|
||||
content_type = content_type_map.get(audio_format, "application/octet-stream")
|
||||
|
||||
try:
|
||||
# 下载临时音频
|
||||
resp = httpx.get(temp_url, timeout=60.0, follow_redirects=True)
|
||||
resp.raise_for_status()
|
||||
audio_data = resp.content
|
||||
|
||||
# 上传到 OSS
|
||||
file_obj = io.BytesIO(audio_data)
|
||||
permanent_url = self._storage.upload_file(file_obj, storage_key, content_type=content_type)
|
||||
logger.info(f"音频转存 OSS 成功: job_id={job_id}, " f"storage_key={storage_key}, size={len(audio_data)}")
|
||||
return permanent_url, storage_key
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"音频转存 OSS 失败,使用临时 URL: " f"job_id={job_id}, error={e}")
|
||||
return temp_url, ""
|
||||
|
||||
def start_synthesis(
|
||||
self,
|
||||
@@ -80,6 +142,10 @@ class TTSWorkflowService:
|
||||
job.mark_processing()
|
||||
job = self.repository.update(job)
|
||||
|
||||
# 长文本自动分段合成
|
||||
if len(job.input_text) > _SEGMENT_THRESHOLD:
|
||||
return self._start_segment_synthesis(job)
|
||||
|
||||
try:
|
||||
submit_result = self.cosyvoice_service.submit_synthesize_task(
|
||||
text=job.input_text,
|
||||
@@ -93,17 +159,19 @@ class TTSWorkflowService:
|
||||
job_metadata["cosyvoice_task_id"] = submit_result.get("task_id", "")
|
||||
job_metadata["cosyvoice_request_id"] = submit_result.get("request_id", "")
|
||||
|
||||
# 如果 CosyVoice 同步返回了 audio_url,直接标记完成
|
||||
# 如果 CosyVoice 同步返回了 audio_url,转存 OSS 后标记完成
|
||||
audio_url = submit_result.get("audio_url", "")
|
||||
if audio_url:
|
||||
permanent_url, storage_key = self._transfer_audio_to_oss(audio_url, job.user_id, job.id, job.format)
|
||||
job.mark_completed(
|
||||
output_audio_url=audio_url,
|
||||
output_audio_url=permanent_url,
|
||||
output_audio_key=storage_key,
|
||||
duration=submit_result.get("duration", 0.0),
|
||||
file_size=submit_result.get("file_size", 0),
|
||||
)
|
||||
job.metadata = job_metadata
|
||||
job = self.repository.update(job)
|
||||
logger.info(f"TTS 合成同步完成: job_id={job.id}, audio_url={audio_url}")
|
||||
logger.info(f"TTS 合成同步完成: job_id={job.id}, audio_url={permanent_url}")
|
||||
return job
|
||||
|
||||
job.metadata = job_metadata
|
||||
@@ -133,6 +201,11 @@ class TTSWorkflowService:
|
||||
if job is None:
|
||||
raise TTSJobNotFoundError(f"TTS job {job_id} not found")
|
||||
|
||||
# 检查是否为分段合成任务
|
||||
segment_task_ids = (job.metadata or {}).get("segment_task_ids", [])
|
||||
if segment_task_ids:
|
||||
return self._poll_segment_tasks(job)
|
||||
|
||||
task_id = (job.metadata or {}).get("cosyvoice_task_id", "")
|
||||
if not task_id:
|
||||
raise ValueError(f"TTSJob {job_id} has no cosyvoice_task_id in metadata")
|
||||
@@ -171,13 +244,17 @@ class TTSWorkflowService:
|
||||
if job is None:
|
||||
raise TTSJobNotFoundError(f"TTS job {job_id} not found")
|
||||
|
||||
# 转存音频到 OSS,获取永久 URL
|
||||
permanent_url, storage_key = self._transfer_audio_to_oss(audio_url, job.user_id, job.id, job.format)
|
||||
|
||||
job.mark_completed(
|
||||
output_audio_url=audio_url,
|
||||
output_audio_url=permanent_url,
|
||||
output_audio_key=storage_key,
|
||||
duration=duration,
|
||||
file_size=file_size,
|
||||
)
|
||||
job = self.repository.update(job)
|
||||
logger.info(f"TTS 合成成功: job_id={job_id}, audio_url={audio_url}")
|
||||
logger.info(f"TTS 合成成功: job_id={job_id}, audio_url={permanent_url}")
|
||||
return job
|
||||
|
||||
def process_synthesis_failure(self, job_id: str, error_message: str) -> TTSJob:
|
||||
@@ -201,3 +278,224 @@ class TTSWorkflowService:
|
||||
job = self.repository.update(job)
|
||||
logger.error(f"TTS 合成失败: job_id={job_id}, error={error_message}")
|
||||
return job
|
||||
|
||||
# ── P1: 长文本分段合成 ─────────────────────────────────────
|
||||
|
||||
def _upload_merged_to_oss(
|
||||
self, merged_data: bytes, user_id: str, job_id: str, audio_format: str
|
||||
) -> tuple[str, str]:
|
||||
"""上传合并后的音频数据到 OSS。
|
||||
|
||||
Returns:
|
||||
(permanent_url, storage_key) 元组。
|
||||
上传失败时返回 ("", "")。
|
||||
"""
|
||||
storage_key = f"tts-outputs/{user_id}/{job_id}.{audio_format}"
|
||||
content_type_map = {
|
||||
"mp3": "audio/mpeg",
|
||||
"wav": "audio/wav",
|
||||
"pcm": "audio/pcm",
|
||||
"opus": "audio/opus",
|
||||
}
|
||||
content_type = content_type_map.get(audio_format, "application/octet-stream")
|
||||
try:
|
||||
file_obj = io.BytesIO(merged_data)
|
||||
permanent_url = self._storage.upload_file(file_obj, storage_key, content_type=content_type)
|
||||
return permanent_url, storage_key
|
||||
except Exception as e:
|
||||
logger.warning(f"分段合并音频转存 OSS 失败: job_id={job_id}, error={e}")
|
||||
return "", ""
|
||||
|
||||
def _start_segment_synthesis(self, job: TTSJob) -> TTSJob:
|
||||
"""长文本分段合成入口。
|
||||
|
||||
将文本分段后并发提交到 CosyVoice,根据同步/异步结果走不同路径。
|
||||
"""
|
||||
segments = split_text(job.input_text, max_chars=_SEGMENT_THRESHOLD)
|
||||
logger.info(f"长文本分段合成: job_id={job.id}, " f"原文={len(job.input_text)}字, 段数={len(segments)}")
|
||||
|
||||
# 记录分段信息到 metadata
|
||||
job_metadata = dict(job.metadata)
|
||||
job_metadata["segment_count"] = len(segments)
|
||||
|
||||
# 并发提交所有分段
|
||||
results = self._submit_segments_concurrent(segments, job)
|
||||
if results is None:
|
||||
# 提交阶段已失败,_submit_segments_concurrent 内部已标记 failed
|
||||
return self.repository.get(job.id)
|
||||
|
||||
# 判断同步还是异步
|
||||
has_audio_urls = any(r.get("audio_url", "") for r in results)
|
||||
has_task_ids = any(r.get("task_id", "") for r in results)
|
||||
|
||||
if has_audio_urls and not has_task_ids:
|
||||
# 所有分段同步返回音频,直接合并
|
||||
return self._process_segments_sync(job, results)
|
||||
|
||||
# 异步路径:保存各分段的 task_id 供后续轮询
|
||||
segment_task_ids = [r.get("task_id", "") for r in results]
|
||||
segment_audio_urls = [r.get("audio_url", "") for r in results]
|
||||
job_metadata["segment_task_ids"] = segment_task_ids
|
||||
job_metadata["segment_audio_urls"] = segment_audio_urls
|
||||
job_metadata["segment_format"] = job.format
|
||||
|
||||
job.metadata = job_metadata
|
||||
job = self.repository.update(job)
|
||||
logger.info(f"分段合成任务已提交(异步): job_id={job.id}, " f"段数={len(segments)}")
|
||||
return job
|
||||
|
||||
def _submit_segments_concurrent(self, segments: list[str], job: TTSJob) -> list[dict] | None:
|
||||
"""并发提交分段合成任务。
|
||||
|
||||
Returns:
|
||||
各分段的结果列表(保持顺序),提交失败时返回 None。
|
||||
"""
|
||||
max_workers = min(len(segments), _MAX_SEGMENT_WORKERS)
|
||||
results: list[dict | None] = [None] * len(segments)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||
future_to_idx = {}
|
||||
for idx, segment_text in enumerate(segments):
|
||||
future = executor.submit(
|
||||
self.cosyvoice_service.submit_synthesize_task,
|
||||
text=segment_text,
|
||||
voice_id=job.voice_id,
|
||||
sample_rate=job.sample_rate,
|
||||
format=job.format,
|
||||
)
|
||||
future_to_idx[future] = idx
|
||||
|
||||
for future in as_completed(future_to_idx):
|
||||
idx = future_to_idx[future]
|
||||
try:
|
||||
results[idx] = future.result()
|
||||
except Exception as e:
|
||||
logger.error(f"分段合成提交失败: job_id={job.id}, " f"segment={idx}, error={e}")
|
||||
self._handle_segment_failure(job, f"分段 {idx + 1} 合成提交失败: {e}")
|
||||
return None
|
||||
|
||||
return results # type: ignore[return-value]
|
||||
|
||||
def _process_segments_sync(self, job: TTSJob, results: list[dict]) -> TTSJob:
|
||||
"""同步路径:所有分段已返回 audio_url,下载合并后转存 OSS。"""
|
||||
merged_data, total_duration = self._download_and_merge_segments(results, job)
|
||||
|
||||
# 直接上传合并后的音频 bytes 到 OSS
|
||||
permanent_url, storage_key = self._upload_merged_to_oss(merged_data, job.user_id, job.id, job.format)
|
||||
|
||||
job.mark_completed(
|
||||
output_audio_url=permanent_url,
|
||||
output_audio_key=storage_key,
|
||||
duration=total_duration,
|
||||
file_size=len(merged_data),
|
||||
)
|
||||
job = self.repository.update(job)
|
||||
logger.info(f"分段合成完成: job_id={job.id}, " f"merged_size={len(merged_data)}, duration={total_duration:.1f}")
|
||||
return job
|
||||
|
||||
def _download_and_merge_segments(self, results: list[dict], job: TTSJob) -> tuple[bytes, float]:
|
||||
"""下载各分段音频并合并。
|
||||
|
||||
Returns:
|
||||
(merged_audio_bytes, total_duration)
|
||||
"""
|
||||
temp_dir = tempfile.mkdtemp(prefix="tts_segments_")
|
||||
try:
|
||||
audio_paths: list[str] = []
|
||||
total_duration = 0.0
|
||||
|
||||
for idx, result in enumerate(results):
|
||||
audio_url = result.get("audio_url", "")
|
||||
if not audio_url:
|
||||
raise TTSWorkflowError(f"分段 {idx + 1} 没有返回 audio_url")
|
||||
|
||||
total_duration += result.get("duration", 0.0)
|
||||
|
||||
# 下载分段音频到临时文件
|
||||
resp = httpx.get(audio_url, timeout=60.0, follow_redirects=True)
|
||||
resp.raise_for_status()
|
||||
|
||||
seg_path = os.path.join(temp_dir, f"seg_{idx:03d}.{job.format}")
|
||||
with open(seg_path, "wb") as f:
|
||||
f.write(resp.content)
|
||||
audio_paths.append(seg_path)
|
||||
|
||||
# 合并
|
||||
merger = AudioMerger()
|
||||
merged_data = merger.merge(audio_paths, output_format=job.format)
|
||||
return merged_data, total_duration
|
||||
|
||||
finally:
|
||||
shutil.rmtree(temp_dir, ignore_errors=True)
|
||||
|
||||
def _poll_segment_tasks(self, job: TTSJob) -> TTSJob:
|
||||
"""轮询所有分段异步任务,全部完成后合并音频。"""
|
||||
segment_task_ids: list[str] = (job.metadata or {}).get("segment_task_ids", [])
|
||||
segment_audio_urls: list[str] = (job.metadata or {}).get("segment_audio_urls", [])
|
||||
segment_count = len(segment_task_ids)
|
||||
|
||||
poll_start = time.monotonic()
|
||||
poll_timeout = 300.0 # 分段任务超时更长
|
||||
poll_interval = 2.0
|
||||
|
||||
while time.monotonic() - poll_start < poll_timeout:
|
||||
all_done = True
|
||||
results: list[dict | None] = [None] * segment_count
|
||||
|
||||
for idx, task_id in enumerate(segment_task_ids):
|
||||
# 已经有音频的分段跳过轮询
|
||||
if idx < len(segment_audio_urls) and segment_audio_urls[idx]:
|
||||
results[idx] = {
|
||||
"audio_url": segment_audio_urls[idx],
|
||||
"duration": 0.0,
|
||||
"file_size": 0,
|
||||
}
|
||||
continue
|
||||
|
||||
try:
|
||||
result = self.cosyvoice_service.poll_synthesize_task(task_id, timeout=poll_timeout)
|
||||
results[idx] = result
|
||||
except Exception as e:
|
||||
logger.error(f"分段任务轮询失败: job_id={job.id}, " f"segment={idx}, error={e}")
|
||||
self._handle_segment_failure(job, f"分段 {idx + 1} 轮询失败: {e}")
|
||||
return self.repository.get(job.id)
|
||||
|
||||
if results[idx] is None:
|
||||
all_done = False
|
||||
|
||||
if all_done and all(r is not None for r in results):
|
||||
# 所有分段完成,下载合并
|
||||
try:
|
||||
merged_data, total_duration = self._download_and_merge_segments(results, job)
|
||||
|
||||
# 转存 OSS
|
||||
permanent_url, storage_key = self._upload_merged_to_oss(
|
||||
merged_data, job.user_id, job.id, job.format
|
||||
)
|
||||
|
||||
job.mark_completed(
|
||||
output_audio_url=permanent_url,
|
||||
output_audio_key=storage_key,
|
||||
duration=total_duration,
|
||||
file_size=len(merged_data),
|
||||
)
|
||||
job = self.repository.update(job)
|
||||
logger.info(f"分段合成轮询完成: job_id={job.id}, " f"merged_size={len(merged_data)}")
|
||||
return job
|
||||
|
||||
except Exception as e:
|
||||
self._handle_segment_failure(job, f"分段合并失败: {e}")
|
||||
return self.repository.get(job.id)
|
||||
|
||||
# 等待后重试
|
||||
time.sleep(poll_interval)
|
||||
|
||||
# 超时
|
||||
self._handle_segment_failure(job, "分段合成轮询超时(300 秒)")
|
||||
return self.repository.get(job.id)
|
||||
|
||||
def _handle_segment_failure(self, job: TTSJob, error_message: str) -> None:
|
||||
"""分段合成失败处理。"""
|
||||
job.mark_failed(error_message)
|
||||
self.repository.update(job)
|
||||
logger.error(f"分段合成失败: job_id={job.id}, error={error_message}")
|
||||
|
||||
@@ -24,6 +24,7 @@ from .entities import (
|
||||
from .generated_video import GeneratedVideo
|
||||
from .generation_task import GenerationTask, GenerationTaskStatus
|
||||
from .job import Job, JobStatus, JobType
|
||||
from .tag import Tag
|
||||
from .template_clip_config import ClipType, TemplateClipConfig, TransitionEffect
|
||||
from .title_library import TitleLibraryItem
|
||||
from .voice_library import VoiceLibraryItem
|
||||
@@ -56,6 +57,7 @@ __all__ = [
|
||||
"JobStatus",
|
||||
"JobType",
|
||||
"Project",
|
||||
"Tag",
|
||||
"TemplateClipConfig",
|
||||
"TransitionEffect",
|
||||
"User",
|
||||
|
||||
+14
-14
@@ -162,7 +162,7 @@ class Asset:
|
||||
quality_score: float | None = None
|
||||
uploaded_by_user_id: str = ""
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
tags: list[str] = field(default_factory=list)
|
||||
tag_ids: list[str] = field(default_factory=list)
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@@ -214,23 +214,23 @@ class Asset:
|
||||
quality_score=quality_score,
|
||||
uploaded_by_user_id=uploaded_by_user_id.strip(),
|
||||
metadata=metadata or {},
|
||||
tags=[],
|
||||
tag_ids=[],
|
||||
)
|
||||
|
||||
def add_tag(self, tag: str) -> None:
|
||||
"""添加标签。空标签会被忽略,自动去重。"""
|
||||
clean_tag = tag.strip()
|
||||
if not clean_tag:
|
||||
raise ValueError("标签不能为空")
|
||||
if clean_tag not in self.tags:
|
||||
self.tags.append(clean_tag)
|
||||
def add_tag(self, tag_id: str) -> None:
|
||||
"""添加标签 ID。空 ID 会被忽略,自动去重。"""
|
||||
clean_id = tag_id.strip()
|
||||
if not clean_id:
|
||||
raise ValueError("标签 ID 不能为空")
|
||||
if clean_id not in self.tag_ids:
|
||||
self.tag_ids.append(clean_id)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def remove_tag(self, tag: str) -> None:
|
||||
"""删除标签。如果标签不存在,不报错(幂等性)。"""
|
||||
clean_tag = tag.strip()
|
||||
if clean_tag in self.tags:
|
||||
self.tags.remove(clean_tag)
|
||||
def remove_tag(self, tag_id: str) -> None:
|
||||
"""删除标签 ID。如果标签不存在,不报错(幂等性)。"""
|
||||
clean_id = tag_id.strip()
|
||||
if clean_id in self.tag_ids:
|
||||
self.tag_ids.remove(clean_id)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
"""标签领域实体。"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class Tag:
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@classmethod
|
||||
def create(cls, user_id: str, name: str) -> "Tag":
|
||||
clean_name = name.strip()
|
||||
if not clean_name:
|
||||
raise ValueError("标签名称不能为空")
|
||||
return cls(
|
||||
id=uuid4().hex,
|
||||
user_id=user_id,
|
||||
name=clean_name,
|
||||
)
|
||||
@@ -4,6 +4,7 @@ from .asset_library_repository import AssetLibraryRepository
|
||||
from .asset_repository import AssetRepository
|
||||
from .ingest_job_repository import IngestJobRepository
|
||||
from .project_repository import ProjectRepository
|
||||
from .tag_repository import TagRepository
|
||||
from .title_library_repository import TitleLibraryRepository
|
||||
from .voice_library_repository import VoiceLibraryRepository
|
||||
|
||||
@@ -12,6 +13,7 @@ __all__ = [
|
||||
"AssetRepository",
|
||||
"IngestJobRepository",
|
||||
"ProjectRepository",
|
||||
"TagRepository",
|
||||
"TitleLibraryRepository",
|
||||
"VoiceLibraryRepository",
|
||||
]
|
||||
|
||||
@@ -50,6 +50,11 @@ class AssetRepository(ABC):
|
||||
def delete(self, asset_id: str) -> bool:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def batch_delete(self, asset_ids: list[str]) -> int:
|
||||
"""批量删除素材,返回实际删除数量。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def count_by_project(self, project_id: str) -> int:
|
||||
pass
|
||||
@@ -78,3 +83,13 @@ class AssetRepository(ABC):
|
||||
) -> list[Asset]:
|
||||
"""按筛选条件搜索候选素材,按质量分降序排列。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_by_tag_ids(
|
||||
self,
|
||||
tag_ids: list[str],
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
) -> list[Asset]:
|
||||
"""查找包含所有指定标签的素材。"""
|
||||
pass
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
"""标签仓储接口定义。"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from packages.domain import Tag
|
||||
|
||||
|
||||
class TagRepository(ABC):
|
||||
@abstractmethod
|
||||
def create(self, tag: Tag) -> Tag:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get(self, tag_id: str) -> Tag | None:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_by_name(self, user_id: str, name: str) -> Tag | None:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def list_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
) -> list[Tag]:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def count_by_user(self, user_id: str) -> int:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def delete(self, tag_id: str) -> bool:
|
||||
pass
|
||||
@@ -4,7 +4,7 @@ from packages.domain import Asset
|
||||
|
||||
|
||||
def test_add_tag_to_asset():
|
||||
"""测试添加标签到 Asset。"""
|
||||
"""测试添加标签 ID 到 Asset。"""
|
||||
asset = Asset.create(
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
@@ -13,16 +13,16 @@ def test_add_tag_to_asset():
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
|
||||
asset.add_tag("风景")
|
||||
asset.add_tag("自然")
|
||||
asset.add_tag("tag-1")
|
||||
asset.add_tag("tag-2")
|
||||
|
||||
assert len(asset.tags) == 2
|
||||
assert "风景" in asset.tags
|
||||
assert "自然" in asset.tags
|
||||
assert len(asset.tag_ids) == 2
|
||||
assert "tag-1" in asset.tag_ids
|
||||
assert "tag-2" in asset.tag_ids
|
||||
|
||||
|
||||
def test_add_duplicate_tag_should_ignore():
|
||||
"""测试添加重复标签应自动去重。"""
|
||||
"""测试添加重复标签 ID 应自动去重。"""
|
||||
asset = Asset.create(
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
@@ -31,15 +31,15 @@ def test_add_duplicate_tag_should_ignore():
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
|
||||
asset.add_tag("风景")
|
||||
asset.add_tag("风景") # 重复
|
||||
asset.add_tag("tag-1")
|
||||
asset.add_tag("tag-1") # 重复
|
||||
|
||||
assert len(asset.tags) == 1
|
||||
assert asset.tags.count("风景") == 1
|
||||
assert len(asset.tag_ids) == 1
|
||||
assert asset.tag_ids.count("tag-1") == 1
|
||||
|
||||
|
||||
def test_add_empty_tag_should_fail():
|
||||
"""测试添加空标签应失败。"""
|
||||
"""测试添加空标签 ID 应失败。"""
|
||||
asset = Asset.create(
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
@@ -48,15 +48,15 @@ def test_add_empty_tag_should_fail():
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="标签不能为空"):
|
||||
with pytest.raises(ValueError, match="标签 ID 不能为空"):
|
||||
asset.add_tag("")
|
||||
|
||||
with pytest.raises(ValueError, match="标签不能为空"):
|
||||
with pytest.raises(ValueError, match="标签 ID 不能为空"):
|
||||
asset.add_tag(" ") # 仅空格
|
||||
|
||||
|
||||
def test_remove_tag_from_asset():
|
||||
"""测试从 Asset 删除标签。"""
|
||||
"""测试从 Asset 删除标签 ID。"""
|
||||
asset = Asset.create(
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
@@ -65,18 +65,18 @@ def test_remove_tag_from_asset():
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
|
||||
asset.add_tag("风景")
|
||||
asset.add_tag("自然")
|
||||
asset.add_tag("tag-1")
|
||||
asset.add_tag("tag-2")
|
||||
|
||||
asset.remove_tag("风景")
|
||||
asset.remove_tag("tag-1")
|
||||
|
||||
assert len(asset.tags) == 1
|
||||
assert "风景" not in asset.tags
|
||||
assert "自然" in asset.tags
|
||||
assert len(asset.tag_ids) == 1
|
||||
assert "tag-1" not in asset.tag_ids
|
||||
assert "tag-2" in asset.tag_ids
|
||||
|
||||
|
||||
def test_remove_nonexistent_tag_should_be_idempotent():
|
||||
"""测试删除不存在的标签应幂等(不报错)。"""
|
||||
"""测试删除不存在的标签 ID 应幂等(不报错)。"""
|
||||
asset = Asset.create(
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
@@ -85,10 +85,10 @@ def test_remove_nonexistent_tag_should_be_idempotent():
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
|
||||
asset.add_tag("风景")
|
||||
asset.add_tag("tag-1")
|
||||
|
||||
# 删除不存在的标签,不应报错
|
||||
asset.remove_tag("不存在的标签")
|
||||
# 删除不存在的标签 ID,不应报错
|
||||
asset.remove_tag("nonexistent-tag")
|
||||
|
||||
assert len(asset.tags) == 1
|
||||
assert "风景" in asset.tags
|
||||
assert len(asset.tag_ids) == 1
|
||||
assert "tag-1" in asset.tag_ids
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
"""批量删除素材 + 分页优化 单元测试。"""
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from packages.adapters.in_memory.asset_repository import InMemoryAssetRepository
|
||||
from packages.domain import Asset, AssetStatus
|
||||
|
||||
|
||||
class TestBatchDelete:
|
||||
"""batch_delete 仓储方法测试。"""
|
||||
|
||||
def _make_repo_with_assets(self):
|
||||
repo = InMemoryAssetRepository()
|
||||
for i in range(5):
|
||||
asset = Asset.create(
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
name=f"voice_{i}.mp3",
|
||||
storage_key=f"uploads/voice_{i}.mp3",
|
||||
mime_type="audio/mpeg",
|
||||
status=AssetStatus.READY,
|
||||
)
|
||||
repo.create(asset)
|
||||
return repo
|
||||
|
||||
def test_batch_delete_removes_multiple(self):
|
||||
repo = InMemoryAssetRepository()
|
||||
assets = []
|
||||
for i in range(5):
|
||||
asset = Asset.create(
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
name=f"voice_{i}.mp3",
|
||||
storage_key=f"uploads/voice_{i}.mp3",
|
||||
mime_type="audio/mpeg",
|
||||
)
|
||||
repo.create(asset)
|
||||
assets.append(asset)
|
||||
|
||||
ids_to_delete = [assets[0].id, assets[2].id, assets[4].id]
|
||||
deleted_count = repo.batch_delete(ids_to_delete)
|
||||
|
||||
assert deleted_count == 3
|
||||
# 验证确实被删了
|
||||
assert repo.get(assets[0].id) is None
|
||||
assert repo.get(assets[2].id) is None
|
||||
assert repo.get(assets[4].id) is None
|
||||
# 验证其他还在
|
||||
assert repo.get(assets[1].id) is not None
|
||||
assert repo.get(assets[3].id) is not None
|
||||
|
||||
def test_batch_delete_empty_list(self):
|
||||
repo = self._make_repo_with_assets()
|
||||
assert repo.batch_delete([]) == 0
|
||||
|
||||
def test_batch_delete_nonexistent_ids(self):
|
||||
repo = self._make_repo_with_assets()
|
||||
deleted = repo.batch_delete(["nonexistent-1", "nonexistent-2"])
|
||||
assert deleted == 0
|
||||
|
||||
def test_batch_delete_mixed_existing_and_nonexistent(self):
|
||||
repo = InMemoryAssetRepository()
|
||||
asset = Asset.create(
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
name="voice.mp3",
|
||||
storage_key="uploads/voice.mp3",
|
||||
mime_type="audio/mpeg",
|
||||
)
|
||||
repo.create(asset)
|
||||
|
||||
deleted = repo.batch_delete([asset.id, "nonexistent"])
|
||||
assert deleted == 1
|
||||
assert repo.get(asset.id) is None
|
||||
@@ -0,0 +1,130 @@
|
||||
"""素材打标/取消标签单元测试。"""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.adapters.in_memory.asset_repository import InMemoryAssetRepository
|
||||
from packages.adapters.in_memory.tag_repository import InMemoryTagRepository
|
||||
from packages.domain import Asset, Tag
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def asset_repo():
|
||||
return InMemoryAssetRepository()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tag_repo():
|
||||
return InMemoryTagRepository()
|
||||
|
||||
|
||||
def _create_asset(asset_repo, **kwargs):
|
||||
defaults = dict(
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
name="video.mp4",
|
||||
storage_key="uploads/abc/video.mp4",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
asset = Asset.create(**defaults)
|
||||
return asset_repo.create(asset)
|
||||
|
||||
|
||||
def test_tag_asset(asset_repo, tag_repo):
|
||||
"""测试给素材打标签。"""
|
||||
asset = _create_asset(asset_repo)
|
||||
tag = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
|
||||
|
||||
asset.add_tag(tag.id)
|
||||
asset_repo.update(asset)
|
||||
|
||||
loaded = asset_repo.get(asset.id)
|
||||
assert tag.id in loaded.tag_ids
|
||||
|
||||
|
||||
def test_untag_asset(asset_repo, tag_repo):
|
||||
"""测试取消素材标签。"""
|
||||
asset = _create_asset(asset_repo)
|
||||
tag = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
|
||||
|
||||
asset.add_tag(tag.id)
|
||||
asset_repo.update(asset)
|
||||
|
||||
asset.remove_tag(tag.id)
|
||||
asset_repo.update(asset)
|
||||
|
||||
loaded = asset_repo.get(asset.id)
|
||||
assert tag.id not in loaded.tag_ids
|
||||
|
||||
|
||||
def test_tag_multiple_assets(asset_repo, tag_repo):
|
||||
"""测试同一标签打给多个素材。"""
|
||||
a1 = _create_asset(asset_repo, name="a.mp4")
|
||||
a2 = _create_asset(asset_repo, name="b.mp4")
|
||||
tag = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
|
||||
|
||||
a1.add_tag(tag.id)
|
||||
a2.add_tag(tag.id)
|
||||
asset_repo.update(a1)
|
||||
asset_repo.update(a2)
|
||||
|
||||
assert tag.id in asset_repo.get(a1.id).tag_ids
|
||||
assert tag.id in asset_repo.get(a2.id).tag_ids
|
||||
|
||||
|
||||
def test_find_by_tag_ids(asset_repo, tag_repo):
|
||||
"""测试按标签 ID 筛选素材。"""
|
||||
tag1 = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
|
||||
tag2 = tag_repo.create(Tag.create(user_id="user-1", name="自然"))
|
||||
|
||||
a1 = _create_asset(asset_repo, name="a.mp4")
|
||||
a1.add_tag(tag1.id)
|
||||
a1.add_tag(tag2.id)
|
||||
asset_repo.update(a1)
|
||||
|
||||
a2 = _create_asset(asset_repo, name="b.mp4")
|
||||
a2.add_tag(tag1.id)
|
||||
asset_repo.update(a2)
|
||||
|
||||
a3 = _create_asset(asset_repo, name="c.mp4")
|
||||
# 无标签
|
||||
|
||||
# 按 tag1 筛选 → a1, a2
|
||||
result = asset_repo.find_by_tag_ids([tag1.id])
|
||||
ids = {a.id for a in result}
|
||||
assert ids == {a1.id, a2.id}
|
||||
|
||||
# 按 tag1 + tag2 筛选(交集)→ a1
|
||||
result = asset_repo.find_by_tag_ids([tag1.id, tag2.id])
|
||||
ids = {a.id for a in result}
|
||||
assert ids == {a1.id}
|
||||
|
||||
# 空 tag_ids → 空结果
|
||||
assert asset_repo.find_by_tag_ids([]) == []
|
||||
|
||||
|
||||
def test_delete_tag_cleans_associations(asset_repo, tag_repo):
|
||||
"""测试删除标签后素材的 tag_ids 不受影响(关联表由仓储层清理)。"""
|
||||
asset = _create_asset(asset_repo)
|
||||
tag = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
|
||||
|
||||
asset.add_tag(tag.id)
|
||||
asset_repo.update(asset)
|
||||
|
||||
# 删除标签
|
||||
tag_repo.delete(tag.id)
|
||||
assert tag_repo.get(tag.id) is None
|
||||
|
||||
# 素材的 tag_ids 在内存中仍有,但重新加载后 InMemory 不感知关联表
|
||||
# 实际 SQLAlchemy 实现中 _sync_asset_tags 会在 update 时清理
|
||||
|
||||
|
||||
def test_duplicate_tag_id_ignored(asset_repo, tag_repo):
|
||||
"""测试重复打同一标签自动去重。"""
|
||||
asset = _create_asset(asset_repo)
|
||||
tag = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
|
||||
|
||||
asset.add_tag(tag.id)
|
||||
asset.add_tag(tag.id) # 重复
|
||||
|
||||
assert asset.tag_ids.count(tag.id) == 1
|
||||
@@ -0,0 +1,102 @@
|
||||
"""标签 CRUD 单元测试(使用 InMemoryTagRepository)。"""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.adapters.in_memory.tag_repository import InMemoryTagRepository
|
||||
from packages.domain import Tag
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tag_repo():
|
||||
return InMemoryTagRepository()
|
||||
|
||||
|
||||
def test_create_tag(tag_repo):
|
||||
"""测试创建标签。"""
|
||||
tag = Tag.create(user_id="user-1", name="风景")
|
||||
created = tag_repo.create(tag)
|
||||
|
||||
assert created.id == tag.id
|
||||
assert created.user_id == "user-1"
|
||||
assert created.name == "风景"
|
||||
|
||||
|
||||
def test_create_tag_strips_whitespace(tag_repo):
|
||||
"""测试创建标签时自动去除首尾空格。"""
|
||||
tag = Tag.create(user_id="user-1", name=" 风景 ")
|
||||
assert tag.name == "风景"
|
||||
|
||||
|
||||
def test_create_tag_empty_name_raises():
|
||||
"""测试空名称抛出 ValueError。"""
|
||||
with pytest.raises(ValueError, match="标签名称不能为空"):
|
||||
Tag.create(user_id="user-1", name="")
|
||||
|
||||
with pytest.raises(ValueError, match="标签名称不能为空"):
|
||||
Tag.create(user_id="user-1", name=" ")
|
||||
|
||||
|
||||
def test_get_tag(tag_repo):
|
||||
"""测试按 ID 获取标签。"""
|
||||
tag = Tag.create(user_id="user-1", name="风景")
|
||||
tag_repo.create(tag)
|
||||
|
||||
found = tag_repo.get(tag.id)
|
||||
assert found is not None
|
||||
assert found.name == "风景"
|
||||
|
||||
assert tag_repo.get("nonexistent") is None
|
||||
|
||||
|
||||
def test_find_by_name(tag_repo):
|
||||
"""测试按用户 ID + 名称查找标签。"""
|
||||
tag = Tag.create(user_id="user-1", name="风景")
|
||||
tag_repo.create(tag)
|
||||
|
||||
found = tag_repo.find_by_name("user-1", "风景")
|
||||
assert found is not None
|
||||
assert found.id == tag.id
|
||||
|
||||
# 不同用户同名标签不冲突
|
||||
assert tag_repo.find_by_name("user-2", "风景") is None
|
||||
|
||||
# 不存在的名称
|
||||
assert tag_repo.find_by_name("user-1", "不存在") is None
|
||||
|
||||
|
||||
def test_list_by_user(tag_repo):
|
||||
"""测试按用户列出标签(分页)。"""
|
||||
for i in range(5):
|
||||
tag_repo.create(Tag.create(user_id="user-1", name=f"标签{i}"))
|
||||
# 另一个用户的标签
|
||||
tag_repo.create(Tag.create(user_id="user-2", name="其他用户标签"))
|
||||
|
||||
items = tag_repo.list_by_user("user-1")
|
||||
assert len(items) == 5
|
||||
|
||||
# 分页
|
||||
items_page = tag_repo.list_by_user("user-1", skip=2, limit=2)
|
||||
assert len(items_page) == 2
|
||||
|
||||
|
||||
def test_count_by_user(tag_repo):
|
||||
"""测试按用户统计标签数量。"""
|
||||
for i in range(3):
|
||||
tag_repo.create(Tag.create(user_id="user-1", name=f"标签{i}"))
|
||||
tag_repo.create(Tag.create(user_id="user-2", name="其他"))
|
||||
|
||||
assert tag_repo.count_by_user("user-1") == 3
|
||||
assert tag_repo.count_by_user("user-2") == 1
|
||||
assert tag_repo.count_by_user("user-3") == 0
|
||||
|
||||
|
||||
def test_delete_tag(tag_repo):
|
||||
"""测试删除标签。"""
|
||||
tag = Tag.create(user_id="user-1", name="风景")
|
||||
tag_repo.create(tag)
|
||||
|
||||
assert tag_repo.delete(tag.id) is True
|
||||
assert tag_repo.get(tag.id) is None
|
||||
|
||||
# 重复删除返回 False
|
||||
assert tag_repo.delete(tag.id) is False
|
||||
@@ -0,0 +1,299 @@
|
||||
"""TTS 音频转存 OSS 单元测试。
|
||||
|
||||
验证 TTSWorkflowService 在合成完成后将 CosyVoice 临时音频转存到 OSS,
|
||||
存储永久 URL 到 TTSJob.output_audio_url,OSS key 到 output_audio_key。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.cosyvoice_service import CosyVoiceService
|
||||
from packages.application.tts_job.workflow import TTSWorkflowService
|
||||
from packages.domain.tts_job import TTSJob, TTSJobStatus
|
||||
|
||||
|
||||
def _make_job(**kwargs) -> TTSJob:
|
||||
defaults = {
|
||||
"id": "test_job_001",
|
||||
"user_id": "user_001",
|
||||
"input_text": "测试文本",
|
||||
"voice_id": "voice_001",
|
||||
"voice_model": "",
|
||||
"project_id": "",
|
||||
"voice_clone_profile_id": "",
|
||||
"status": TTSJobStatus.PENDING,
|
||||
"output_audio_url": "",
|
||||
"output_audio_key": "",
|
||||
"duration": 0.0,
|
||||
"file_size": 0,
|
||||
"sample_rate": 22050,
|
||||
"format": "mp3",
|
||||
"error_message": "",
|
||||
"retry_count": 0,
|
||||
"max_retries": 3,
|
||||
"metadata": {},
|
||||
"started_at": None,
|
||||
"completed_at": None,
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
return TTSJob(**defaults)
|
||||
|
||||
|
||||
def _make_workflow(
|
||||
cosyvoice_service: MagicMock | None = None,
|
||||
repo: MagicMock | None = None,
|
||||
storage: MagicMock | None = None,
|
||||
) -> TTSWorkflowService:
|
||||
if cosyvoice_service is None:
|
||||
cosyvoice_service = MagicMock(spec=CosyVoiceService)
|
||||
if repo is None:
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = _make_job()
|
||||
repo.update.side_effect = lambda j: j
|
||||
if storage is None:
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/tts-outputs/user_001/test_job_001.mp3"
|
||||
return TTSWorkflowService(
|
||||
repository=repo,
|
||||
cosyvoice_service=cosyvoice_service,
|
||||
storage_service=storage,
|
||||
)
|
||||
|
||||
|
||||
class TestTransferAudioToOSS:
|
||||
"""测试 _transfer_audio_to_oss 方法。"""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_success_download_and_upload(self, mock_httpx: MagicMock) -> None:
|
||||
"""成功下载音频并上传到 OSS,返回永久 URL 和 storage_key。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"fake audio data"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/tts-outputs/user_001/job_123.mp3"
|
||||
|
||||
workflow = _make_workflow(storage=storage)
|
||||
url, key = workflow._transfer_audio_to_oss(
|
||||
"https://cosyvoice-temp.com/audio.mp3",
|
||||
"user_001",
|
||||
"job_123",
|
||||
"mp3",
|
||||
)
|
||||
|
||||
assert url == "https://oss.example.com/tts-outputs/user_001/job_123.mp3"
|
||||
assert key == "tts-outputs/user_001/job_123.mp3"
|
||||
|
||||
mock_httpx.get.assert_called_once_with(
|
||||
"https://cosyvoice-temp.com/audio.mp3",
|
||||
timeout=60.0,
|
||||
follow_redirects=True,
|
||||
)
|
||||
storage.upload_file.assert_called_once()
|
||||
call_args = storage.upload_file.call_args
|
||||
assert call_args[0][1] == "tts-outputs/user_001/job_123.mp3"
|
||||
assert call_args[1]["content_type"] == "audio/mpeg"
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_download_failure_fallback(self, mock_httpx: MagicMock) -> None:
|
||||
"""下载失败时回退到原始临时 URL,storage_key 为空。"""
|
||||
mock_httpx.get.side_effect = Exception("Network error")
|
||||
|
||||
workflow = _make_workflow()
|
||||
url, key = workflow._transfer_audio_to_oss(
|
||||
"https://cosyvoice-temp.com/audio.mp3",
|
||||
"user_001",
|
||||
"job_123",
|
||||
)
|
||||
|
||||
assert url == "https://cosyvoice-temp.com/audio.mp3"
|
||||
assert key == ""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_upload_failure_fallback(self, mock_httpx: MagicMock) -> None:
|
||||
"""上传 OSS 失败时回退到原始临时 URL。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"fake audio data"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.side_effect = Exception("OSS bucket error")
|
||||
|
||||
workflow = _make_workflow(storage=storage)
|
||||
url, key = workflow._transfer_audio_to_oss(
|
||||
"https://cosyvoice-temp.com/audio.mp3",
|
||||
"user_001",
|
||||
"job_123",
|
||||
)
|
||||
|
||||
assert url == "https://cosyvoice-temp.com/audio.mp3"
|
||||
assert key == ""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_wav_content_type(self, mock_httpx: MagicMock) -> None:
|
||||
"""wav 格式使用正确的 content_type。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"fake wav data"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/audio.wav"
|
||||
|
||||
workflow = _make_workflow(storage=storage)
|
||||
workflow._transfer_audio_to_oss(
|
||||
"https://cosyvoice-temp.com/audio.wav",
|
||||
"user_001",
|
||||
"job_456",
|
||||
"wav",
|
||||
)
|
||||
|
||||
call_args = storage.upload_file.call_args
|
||||
assert call_args[1]["content_type"] == "audio/wav"
|
||||
assert call_args[0][1] == "tts-outputs/user_001/job_456.wav"
|
||||
|
||||
|
||||
class TestProcessSynthesisResultWithOSS:
|
||||
"""测试 process_synthesis_result 集成 OSS 转存。"""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_stores_permanent_url_and_key(self, mock_httpx: MagicMock) -> None:
|
||||
"""合成结果存 OSS 永久 URL 和 storage_key。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"audio bytes"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/tts-outputs/user_001/test_job_001.mp3"
|
||||
|
||||
repo = MagicMock()
|
||||
job = _make_job(status=TTSJobStatus.PROCESSING)
|
||||
repo.get.return_value = job
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(repo=repo, storage=storage)
|
||||
result = workflow.process_synthesis_result(
|
||||
"test_job_001",
|
||||
audio_url="https://cosyvoice-temp.com/expiring.mp3",
|
||||
duration=5.0,
|
||||
file_size=50000,
|
||||
)
|
||||
|
||||
assert result.status == TTSJobStatus.COMPLETED
|
||||
assert result.output_audio_url == "https://oss.example.com/tts-outputs/user_001/test_job_001.mp3"
|
||||
assert result.output_audio_key == "tts-outputs/user_001/test_job_001.mp3"
|
||||
assert result.duration == 5.0
|
||||
assert result.file_size == 50000
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_fallback_to_temp_url_on_oss_failure(self, mock_httpx: MagicMock) -> None:
|
||||
"""OSS 转存失败时,使用 CosyVoice 临时 URL(不阻塞合成流程)。"""
|
||||
mock_httpx.get.side_effect = Exception("Download failed")
|
||||
|
||||
repo = MagicMock()
|
||||
job = _make_job(status=TTSJobStatus.PROCESSING)
|
||||
repo.get.return_value = job
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(repo=repo)
|
||||
result = workflow.process_synthesis_result(
|
||||
"test_job_001",
|
||||
audio_url="https://cosyvoice-temp.com/expiring.mp3",
|
||||
)
|
||||
|
||||
assert result.status == TTSJobStatus.COMPLETED
|
||||
assert result.output_audio_url == "https://cosyvoice-temp.com/expiring.mp3"
|
||||
assert result.output_audio_key == ""
|
||||
|
||||
|
||||
class TestStartSynthesisSyncWithOSS:
|
||||
"""测试 start_synthesis 同步路径的 OSS 转存。"""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_sync_path_transfers_to_oss(self, mock_httpx: MagicMock) -> None:
|
||||
"""CosyVoice 同步返回 audio_url 时,也走 OSS 转存。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"sync audio bytes"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/tts-outputs/user_001/test_job_001.mp3"
|
||||
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
service.submit_synthesize_task.return_value = {
|
||||
"task_id": "",
|
||||
"audio_url": "https://cosyvoice-temp.com/sync.mp3",
|
||||
"duration": 2.0,
|
||||
"file_size": 20000,
|
||||
"request_id": "req_sync",
|
||||
}
|
||||
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = _make_job()
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo, storage=storage)
|
||||
job = workflow.start_synthesis("test_job_001")
|
||||
|
||||
assert job.status == TTSJobStatus.COMPLETED
|
||||
assert job.output_audio_url == "https://oss.example.com/tts-outputs/user_001/test_job_001.mp3"
|
||||
assert job.output_audio_key == "tts-outputs/user_001/test_job_001.mp3"
|
||||
assert job.duration == 2.0
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_sync_path_oss_failure_stores_temp_url(self, mock_httpx: MagicMock) -> None:
|
||||
"""同步路径 OSS 失败时,降级存储临时 URL。"""
|
||||
mock_httpx.get.side_effect = Exception("Network error")
|
||||
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
service.submit_synthesize_task.return_value = {
|
||||
"task_id": "",
|
||||
"audio_url": "https://cosyvoice-temp.com/sync.mp3",
|
||||
"duration": 2.0,
|
||||
"file_size": 20000,
|
||||
"request_id": "req_sync",
|
||||
}
|
||||
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = _make_job()
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo)
|
||||
job = workflow.start_synthesis("test_job_001")
|
||||
|
||||
assert job.status == TTSJobStatus.COMPLETED
|
||||
assert job.output_audio_url == "https://cosyvoice-temp.com/sync.mp3"
|
||||
assert job.output_audio_key == ""
|
||||
|
||||
def test_async_path_no_oss_transfer(self) -> None:
|
||||
"""异步路径(返回 task_id,无 audio_url)不触发 OSS 转存。"""
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
service.submit_synthesize_task.return_value = {
|
||||
"task_id": "cosy_task_async",
|
||||
"audio_url": "",
|
||||
"duration": 0.0,
|
||||
"file_size": 0,
|
||||
"request_id": "req_async",
|
||||
}
|
||||
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = _make_job()
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
storage = MagicMock()
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo, storage=storage)
|
||||
job = workflow.start_synthesis("test_job_001")
|
||||
|
||||
assert job.status == TTSJobStatus.PROCESSING
|
||||
# 异步路径不应调用 OSS 上传
|
||||
storage.upload_file.assert_not_called()
|
||||
@@ -0,0 +1,344 @@
|
||||
"""最后一公里:TTS 合成结果保存到配音库 单元测试。
|
||||
|
||||
覆盖:
|
||||
- 正常保存已完成 TTS job 到配音库
|
||||
- 自动携带元信息(音色名、时长、语速等)
|
||||
- 自定义名称
|
||||
- TTS job 不存在 → 404
|
||||
- TTS job 未完成 → 400
|
||||
- 配音库配额已满 → 429
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.tts_job import TTSJob, TTSJobStatus
|
||||
from packages.domain.voice_library import VoiceLibraryItem
|
||||
|
||||
|
||||
def _make_completed_job(**kwargs) -> TTSJob:
|
||||
"""构造一个已完成的 TTSJob。"""
|
||||
defaults = {
|
||||
"id": "tts_job_001",
|
||||
"user_id": "user_001",
|
||||
"input_text": "你好世界",
|
||||
"voice_id": "voice_001",
|
||||
"voice_model": "CosyVoice-v1",
|
||||
"project_id": "proj_001",
|
||||
"voice_clone_profile_id": "",
|
||||
"status": TTSJobStatus.COMPLETED,
|
||||
"output_audio_url": "https://oss.example.com/audio.mp3",
|
||||
"output_audio_key": "tts-outputs/user_001/tts_job_001.mp3",
|
||||
"duration": 5.5,
|
||||
"file_size": 88000,
|
||||
"sample_rate": 22050,
|
||||
"format": "mp3",
|
||||
"error_message": "",
|
||||
"retry_count": 0,
|
||||
"max_retries": 3,
|
||||
"metadata": {"speed": 1.0, "language": "zh-CN"},
|
||||
"started_at": datetime(2026, 7, 7, 10, 0, 0, tzinfo=timezone.utc),
|
||||
"completed_at": datetime(2026, 7, 7, 10, 0, 5, tzinfo=timezone.utc),
|
||||
"created_at": datetime(2026, 7, 7, 10, 0, 0, tzinfo=timezone.utc),
|
||||
"updated_at": datetime(2026, 7, 7, 10, 0, 5, tzinfo=timezone.utc),
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
return TTSJob(**defaults)
|
||||
|
||||
|
||||
def _make_voice_library_item(**kwargs) -> VoiceLibraryItem:
|
||||
"""构造一个配音库条目。"""
|
||||
defaults = {
|
||||
"id": "voice_lib_001",
|
||||
"user_id": "user_001",
|
||||
"name": "TTS-tts_job_",
|
||||
"text": "你好世界",
|
||||
"voice_provider": "cosyvoice",
|
||||
"voice_id": "voice_001",
|
||||
"voice_name": "CosyVoice-v1",
|
||||
"audio_url": "https://oss.example.com/audio.mp3",
|
||||
"duration": 5.5,
|
||||
"file_size": 88000,
|
||||
"status": "completed",
|
||||
"project_id": "proj_001",
|
||||
"tags": [],
|
||||
"metadata_": {
|
||||
"source": "tts_job",
|
||||
"tts_job_id": "tts_job_001",
|
||||
"format": "mp3",
|
||||
"sample_rate": 22050,
|
||||
"speed": 1.0,
|
||||
"language": "zh-CN",
|
||||
},
|
||||
"created_at": datetime(2026, 7, 7, 10, 1, 0, tzinfo=timezone.utc),
|
||||
"updated_at": datetime(2026, 7, 7, 10, 1, 0, tzinfo=timezone.utc),
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
return VoiceLibraryItem(**defaults)
|
||||
|
||||
|
||||
class TestSaveToLibraryMapping:
|
||||
"""测试 TTSJob → VoiceLibraryItem 字段映射。"""
|
||||
|
||||
def test_completed_job_maps_correctly(self) -> None:
|
||||
"""已完成的 TTS job 字段正确映射到配音库条目。"""
|
||||
job = _make_completed_job()
|
||||
|
||||
# 验证 is_completed 属性
|
||||
assert job.is_completed is True
|
||||
|
||||
# 验证关键字段映射
|
||||
assert job.output_audio_url == "https://oss.example.com/audio.mp3"
|
||||
assert job.duration == 5.5
|
||||
assert job.file_size == 88000
|
||||
assert job.voice_id == "voice_001"
|
||||
assert job.voice_model == "CosyVoice-v1"
|
||||
assert job.input_text == "你好世界"
|
||||
assert job.format == "mp3"
|
||||
assert job.sample_rate == 22050
|
||||
|
||||
def test_metadata_carries_speed_and_format(self) -> None:
|
||||
"""元信息携带语速、格式等。"""
|
||||
job = _make_completed_job()
|
||||
|
||||
metadata = {
|
||||
"source": "tts_job",
|
||||
"tts_job_id": job.id,
|
||||
"format": job.format,
|
||||
"sample_rate": job.sample_rate,
|
||||
}
|
||||
if job.metadata:
|
||||
for key in ("speed", "language"):
|
||||
if key in job.metadata:
|
||||
metadata[key] = job.metadata[key]
|
||||
|
||||
assert metadata["format"] == "mp3"
|
||||
assert metadata["sample_rate"] == 22050
|
||||
assert metadata["speed"] == 1.0
|
||||
assert metadata["language"] == "zh-CN"
|
||||
assert metadata["source"] == "tts_job"
|
||||
|
||||
def test_name_auto_generated_when_empty(self) -> None:
|
||||
"""未提供名称时自动生成。"""
|
||||
job = _make_completed_job()
|
||||
name = None # 模拟未提供名称
|
||||
generated_name = name or f"TTS-{job.id[:8]}"
|
||||
assert generated_name == "TTS-tts_job_"
|
||||
|
||||
def test_name_uses_custom_when_provided(self) -> None:
|
||||
"""提供自定义名称时使用自定义名称。"""
|
||||
custom_name = "我的配音"
|
||||
generated_name = custom_name or "TTS-fallback"
|
||||
assert generated_name == "我的配音"
|
||||
|
||||
|
||||
class TestSaveToLibraryNotCompleted:
|
||||
"""测试未完成 job 不能保存。"""
|
||||
|
||||
def test_pending_job_not_completed(self) -> None:
|
||||
"""pending 状态的 job 不能保存。"""
|
||||
job = _make_completed_job(status=TTSJobStatus.PENDING)
|
||||
assert job.is_completed is False
|
||||
|
||||
def test_processing_job_not_completed(self) -> None:
|
||||
"""processing 状态的 job 不能保存。"""
|
||||
job = _make_completed_job(status=TTSJobStatus.PROCESSING)
|
||||
assert job.is_completed is False
|
||||
|
||||
def test_failed_job_not_completed(self) -> None:
|
||||
"""failed 状态的 job 不能保存。"""
|
||||
job = _make_completed_job(status=TTSJobStatus.FAILED)
|
||||
assert job.is_completed is False
|
||||
|
||||
def test_completed_without_url_not_completed(self) -> None:
|
||||
"""status=completed 但没有 audio_url 的 job 不算完成。"""
|
||||
job = _make_completed_job(
|
||||
status=TTSJobStatus.COMPLETED,
|
||||
output_audio_url="",
|
||||
)
|
||||
assert job.is_completed is False
|
||||
|
||||
|
||||
class TestSaveToLibraryQuota:
|
||||
"""测试配额检查。"""
|
||||
|
||||
def test_quota_exceeded_raises(self) -> None:
|
||||
"""配音库配额已满时抛出 QuotaExceededError。"""
|
||||
from packages.application.voice_library.use_cases import QuotaExceededError
|
||||
|
||||
error = QuotaExceededError(dimension="max_voiceovers", limit=10, used=10)
|
||||
assert "10/10" in str(error)
|
||||
|
||||
def test_quota_under_limit_passes(self) -> None:
|
||||
"""配额未满时不报错。"""
|
||||
from packages.domain.quota import QuotaDimension, quota_checker
|
||||
|
||||
result = quota_checker.check("free", QuotaDimension.MAX_VOICEOVERS.value, 5)
|
||||
assert result.allowed is True
|
||||
|
||||
|
||||
class TestSaveToLibraryCreateCommand:
|
||||
"""测试 CreateVoiceLibraryCommand 构建。"""
|
||||
|
||||
def test_command_fields_from_tts_job(self) -> None:
|
||||
"""从 TTSJob 构建的 Command 字段正确。"""
|
||||
from packages.application.voice_library.commands import CreateVoiceLibraryCommand
|
||||
|
||||
job = _make_completed_job()
|
||||
metadata_ = {
|
||||
"source": "tts_job",
|
||||
"tts_job_id": job.id,
|
||||
"format": job.format,
|
||||
"sample_rate": job.sample_rate,
|
||||
"speed": 1.0,
|
||||
"language": "zh-CN",
|
||||
}
|
||||
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id=job.user_id,
|
||||
name=f"TTS-{job.id[:8]}",
|
||||
text=job.input_text,
|
||||
voice_provider="cosyvoice",
|
||||
voice_id=job.voice_id,
|
||||
voice_name=job.voice_model,
|
||||
audio_url=job.output_audio_url,
|
||||
duration=job.duration,
|
||||
file_size=job.file_size,
|
||||
status="completed",
|
||||
project_id=job.project_id,
|
||||
tags=[],
|
||||
metadata_=metadata_,
|
||||
)
|
||||
|
||||
assert command.user_id == "user_001"
|
||||
assert command.name == "TTS-tts_job_"
|
||||
assert command.text == "你好世界"
|
||||
assert command.voice_provider == "cosyvoice"
|
||||
assert command.voice_id == "voice_001"
|
||||
assert command.voice_name == "CosyVoice-v1"
|
||||
assert command.audio_url == "https://oss.example.com/audio.mp3"
|
||||
assert command.duration == 5.5
|
||||
assert command.file_size == 88000
|
||||
assert command.status == "completed"
|
||||
assert command.project_id == "proj_001"
|
||||
assert command.metadata_["source"] == "tts_job"
|
||||
|
||||
def test_command_with_empty_project_id(self) -> None:
|
||||
"""project_id 为空时传空字符串。"""
|
||||
from packages.application.voice_library.commands import CreateVoiceLibraryCommand
|
||||
|
||||
job = _make_completed_job(project_id="")
|
||||
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id=job.user_id,
|
||||
name="test",
|
||||
text=job.input_text,
|
||||
voice_provider="cosyvoice",
|
||||
voice_id=job.voice_id,
|
||||
voice_name="",
|
||||
audio_url=job.output_audio_url,
|
||||
duration=job.duration,
|
||||
file_size=job.file_size,
|
||||
status="completed",
|
||||
project_id=job.project_id or "",
|
||||
tags=[],
|
||||
metadata_={},
|
||||
)
|
||||
|
||||
assert command.project_id == ""
|
||||
|
||||
def test_command_voice_name_fallback(self) -> None:
|
||||
"""voice_model 为空时 voice_name 回退为空字符串。"""
|
||||
from packages.application.voice_library.commands import CreateVoiceLibraryCommand
|
||||
|
||||
job = _make_completed_job(voice_model="")
|
||||
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id=job.user_id,
|
||||
name="test",
|
||||
text=job.input_text,
|
||||
voice_provider="cosyvoice",
|
||||
voice_id=job.voice_id,
|
||||
voice_name=job.voice_model or "",
|
||||
audio_url=job.output_audio_url,
|
||||
duration=job.duration,
|
||||
file_size=job.file_size,
|
||||
status="completed",
|
||||
project_id="",
|
||||
tags=[],
|
||||
metadata_={},
|
||||
)
|
||||
|
||||
assert command.voice_name == ""
|
||||
|
||||
|
||||
class TestSaveToLibraryUseCase:
|
||||
"""测试 CreateVoiceLibraryUseCase 调用。"""
|
||||
|
||||
def test_use_case_creates_item(self) -> None:
|
||||
"""UseCase 正确创建配音库条目。"""
|
||||
from packages.application.voice_library.commands import CreateVoiceLibraryCommand
|
||||
from packages.application.voice_library.use_cases import CreateVoiceLibraryUseCase
|
||||
|
||||
repo = MagicMock()
|
||||
repo.count_by_user.return_value = 0 # 配额未满
|
||||
|
||||
expected_item = _make_voice_library_item()
|
||||
repo.create.side_effect = lambda item: item
|
||||
|
||||
use_case = CreateVoiceLibraryUseCase(repo)
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user_001",
|
||||
name="test",
|
||||
text="你好",
|
||||
voice_provider="cosyvoice",
|
||||
voice_id="v1",
|
||||
voice_name="Voice1",
|
||||
audio_url="https://example.com/audio.mp3",
|
||||
duration=3.0,
|
||||
file_size=5000,
|
||||
status="completed",
|
||||
project_id="",
|
||||
tags=[],
|
||||
metadata_={},
|
||||
)
|
||||
|
||||
item = use_case.execute(command, plan_name="free")
|
||||
|
||||
repo.create.assert_called_once()
|
||||
assert item is not None
|
||||
|
||||
def test_use_case_quota_exceeded(self) -> None:
|
||||
"""UseCase 配额已满时抛出 QuotaExceededError。"""
|
||||
from packages.application.voice_library.commands import CreateVoiceLibraryCommand
|
||||
from packages.application.voice_library.use_cases import (
|
||||
CreateVoiceLibraryUseCase,
|
||||
QuotaExceededError,
|
||||
)
|
||||
|
||||
repo = MagicMock()
|
||||
repo.count_by_user.return_value = 100 # 超过 premium 配额
|
||||
|
||||
use_case = CreateVoiceLibraryUseCase(repo)
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user_001",
|
||||
name="test",
|
||||
text="你好",
|
||||
voice_provider="cosyvoice",
|
||||
voice_id="v1",
|
||||
voice_name="Voice1",
|
||||
audio_url="https://example.com/audio.mp3",
|
||||
duration=3.0,
|
||||
file_size=5000,
|
||||
status="completed",
|
||||
project_id="",
|
||||
tags=[],
|
||||
metadata_={},
|
||||
)
|
||||
|
||||
with pytest.raises(QuotaExceededError):
|
||||
use_case.execute(command, plan_name="premium")
|
||||
@@ -0,0 +1,485 @@
|
||||
"""P1 长文本分段合成单元测试。
|
||||
|
||||
覆盖:
|
||||
- text_splitter.split_text 分段逻辑
|
||||
- audio_merger.AudioMerger 合并逻辑
|
||||
- workflow 分段合成路径(同步 / 异步 / 失败)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, call, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
|
||||
from packages.application.tts_job.audio_merger import AudioMergeError, AudioMerger
|
||||
from packages.application.tts_job.text_splitter import split_text
|
||||
from packages.application.tts_job.workflow import TTSWorkflowService
|
||||
from packages.domain.tts_job import TTSJob, TTSJobStatus
|
||||
|
||||
# ── text_splitter ─────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSplitText:
|
||||
"""测试文本分段工具。"""
|
||||
|
||||
def test_short_text_no_split(self) -> None:
|
||||
"""短文本不拆分。"""
|
||||
assert split_text("你好世界", max_chars=500) == ["你好世界"]
|
||||
|
||||
def test_empty_text(self) -> None:
|
||||
"""空文本返回空列表。"""
|
||||
assert split_text("") == []
|
||||
assert split_text(" ") == []
|
||||
|
||||
def test_exact_threshold(self) -> None:
|
||||
"""恰好等于阈值不拆分。"""
|
||||
text = "a" * 500
|
||||
assert split_text(text, max_chars=500) == [text]
|
||||
|
||||
def test_split_at_sentence_boundary(self) -> None:
|
||||
"""在句子边界处分段。"""
|
||||
text = "第一句话。" * 60 # 300 chars
|
||||
text += "第二句话。" * 60 # 300 chars → total 600
|
||||
segments = split_text(text, max_chars=500)
|
||||
assert len(segments) >= 2
|
||||
for seg in segments:
|
||||
assert len(seg) <= 500
|
||||
|
||||
def test_split_at_newline(self) -> None:
|
||||
"""在换行符处分段。"""
|
||||
text = "段落一\n" * 100 # 300 chars
|
||||
text += "段落二\n" * 100 # 300 chars
|
||||
segments = split_text(text, max_chars=500)
|
||||
assert len(segments) >= 2
|
||||
|
||||
def test_long_sentence_hard_split(self) -> None:
|
||||
"""超长句子硬切。"""
|
||||
text = "a" * 1200
|
||||
segments = split_text(text, max_chars=500)
|
||||
assert len(segments) >= 3
|
||||
for seg in segments:
|
||||
assert len(seg) <= 500
|
||||
|
||||
def test_merge_short_segments(self) -> None:
|
||||
"""短段合并减少 API 调用。"""
|
||||
# 多个短句子应该被合并
|
||||
text = "你好。" * 120 # 360 chars, each sentence 3 chars
|
||||
segments = split_text(text, max_chars=500)
|
||||
# 短段应该被合并,段数应该比较少
|
||||
assert len(segments) < 120
|
||||
|
||||
def test_preserves_order(self) -> None:
|
||||
"""分段保持原始顺序。"""
|
||||
text = "第一段。第二段。第三段。" + "x" * 490
|
||||
segments = split_text(text, max_chars=500)
|
||||
# 第一个段应该以 "第一段" 开头
|
||||
assert segments[0].startswith("第一段")
|
||||
|
||||
|
||||
# ── audio_merger ─────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAudioMerger:
|
||||
"""测试 FFmpeg 音频合并器。"""
|
||||
|
||||
def test_empty_list_raises(self) -> None:
|
||||
"""空列表抛出 AudioMergeError。"""
|
||||
merger = AudioMerger()
|
||||
with pytest.raises(AudioMergeError, match="没有可合并"):
|
||||
merger.merge([])
|
||||
|
||||
def test_single_file_returns_bytes(self) -> None:
|
||||
"""单文件直接返回内容。"""
|
||||
with tempfile.NamedTemporaryFile(suffix=".mp3", delete=False) as f:
|
||||
f.write(b"fake audio content")
|
||||
f.flush()
|
||||
path = f.name
|
||||
|
||||
try:
|
||||
merger = AudioMerger()
|
||||
data = merger.merge([path])
|
||||
assert data == b"fake audio content"
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
@patch("packages.application.tts_job.audio_merger.subprocess.run")
|
||||
def test_ffmpeg_called_correctly(self, mock_run: MagicMock) -> None:
|
||||
"""多文件调用 FFmpeg concat。"""
|
||||
mock_run.return_value = MagicMock(returncode=0)
|
||||
|
||||
# 创建临时文件
|
||||
paths = []
|
||||
for i in range(3):
|
||||
with tempfile.NamedTemporaryFile(suffix=".mp3", delete=False) as f:
|
||||
f.write(b"audio")
|
||||
paths.append(f.name)
|
||||
|
||||
try:
|
||||
merger = AudioMerger()
|
||||
# Mock open for reading the merged output
|
||||
with patch("builtins.open", create=True) as mock_open:
|
||||
mock_open.return_value.__enter__ = lambda s: s
|
||||
mock_open.return_value.read = lambda: b"merged audio"
|
||||
try:
|
||||
merger.merge(paths, output_format="mp3")
|
||||
except (FileNotFoundError, OSError):
|
||||
pass # Expected since we're mocking
|
||||
|
||||
# 验证 FFmpeg 被调用
|
||||
mock_run.assert_called_once()
|
||||
cmd = mock_run.call_args[0][0]
|
||||
assert cmd[0] == "ffmpeg"
|
||||
assert "-f" in cmd
|
||||
assert "concat" in cmd
|
||||
finally:
|
||||
for p in paths:
|
||||
os.unlink(p)
|
||||
|
||||
@patch("packages.application.tts_job.audio_merger.subprocess.run")
|
||||
def test_ffmpeg_failure_raises(self, mock_run: MagicMock) -> None:
|
||||
"""FFmpeg 失败抛出 AudioMergeError。"""
|
||||
mock_run.return_value = MagicMock(returncode=1, stderr="error details")
|
||||
|
||||
paths = []
|
||||
for i in range(2):
|
||||
with tempfile.NamedTemporaryFile(suffix=".mp3", delete=False) as f:
|
||||
f.write(b"audio")
|
||||
paths.append(f.name)
|
||||
|
||||
try:
|
||||
merger = AudioMerger()
|
||||
with pytest.raises(AudioMergeError, match="FFmpeg 合并失败"):
|
||||
merger.merge(paths)
|
||||
finally:
|
||||
for p in paths:
|
||||
os.unlink(p)
|
||||
|
||||
|
||||
# ── workflow segment methods ─────────────────────────────────
|
||||
|
||||
|
||||
def _make_job(**kwargs) -> TTSJob:
|
||||
defaults = {
|
||||
"id": "test_job_seg",
|
||||
"user_id": "user_001",
|
||||
"input_text": "x" * 600, # > 500 threshold
|
||||
"voice_id": "voice_001",
|
||||
"voice_model": "",
|
||||
"project_id": "",
|
||||
"voice_clone_profile_id": "",
|
||||
"status": TTSJobStatus.PENDING,
|
||||
"output_audio_url": "",
|
||||
"output_audio_key": "",
|
||||
"duration": 0.0,
|
||||
"file_size": 0,
|
||||
"sample_rate": 22050,
|
||||
"format": "mp3",
|
||||
"error_message": "",
|
||||
"retry_count": 0,
|
||||
"max_retries": 3,
|
||||
"metadata": {},
|
||||
"started_at": None,
|
||||
"completed_at": None,
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
return TTSJob(**defaults)
|
||||
|
||||
|
||||
def _make_workflow(
|
||||
cosyvoice_service: MagicMock | None = None,
|
||||
repo: MagicMock | None = None,
|
||||
storage: MagicMock | None = None,
|
||||
) -> TTSWorkflowService:
|
||||
if cosyvoice_service is None:
|
||||
cosyvoice_service = MagicMock(spec=CosyVoiceService)
|
||||
if repo is None:
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = _make_job()
|
||||
repo.update.side_effect = lambda j: j
|
||||
if storage is None:
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/merged.mp3"
|
||||
return TTSWorkflowService(
|
||||
repository=repo,
|
||||
cosyvoice_service=cosyvoice_service,
|
||||
storage_service=storage,
|
||||
)
|
||||
|
||||
|
||||
class TestStartSegmentSynthesis:
|
||||
"""测试 _start_segment_synthesis 分段合成入口。"""
|
||||
|
||||
def test_short_text_no_segment(self) -> None:
|
||||
"""短文本不触发分段。"""
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
service.submit_synthesize_task.return_value = {
|
||||
"task_id": "task_1",
|
||||
"audio_url": "",
|
||||
}
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = _make_job(input_text="短文本")
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo)
|
||||
job = workflow.start_synthesis("test_job_seg")
|
||||
|
||||
# 短文本走普通路径,不调用分段
|
||||
assert job.status == TTSJobStatus.PROCESSING
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_long_text_sync_segments(self, mock_httpx: MagicMock) -> None:
|
||||
"""长文本同步分段:所有段立即返回 audio_url,直接合并。"""
|
||||
# Mock 分段音频下载
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"segment audio"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
# 每个分段都同步返回 audio_url
|
||||
service.submit_synthesize_task.side_effect = [
|
||||
{"task_id": "", "audio_url": "https://temp.com/seg1.mp3", "duration": 2.0, "file_size": 1000},
|
||||
{"task_id": "", "audio_url": "https://temp.com/seg2.mp3", "duration": 3.0, "file_size": 1500},
|
||||
]
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/merged.mp3"
|
||||
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = _make_job()
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo, storage=storage)
|
||||
|
||||
# Mock AudioMerger 避免真实 FFmpeg 调用
|
||||
with patch("packages.application.tts_job.workflow.AudioMerger") as MockMerger:
|
||||
mock_merger = MagicMock()
|
||||
mock_merger.merge.return_value = b"merged audio data"
|
||||
MockMerger.return_value = mock_merger
|
||||
|
||||
job = workflow.start_synthesis("test_job_seg")
|
||||
|
||||
assert job.status == TTSJobStatus.COMPLETED
|
||||
assert job.output_audio_url == "https://oss.example.com/merged.mp3"
|
||||
assert job.duration == 5.0 # 2.0 + 3.0
|
||||
|
||||
def test_long_text_async_segments(self) -> None:
|
||||
"""长文本异步分段:返回 task_id,存入 metadata。"""
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
# 每个分段返回 task_id(异步)
|
||||
service.submit_synthesize_task.side_effect = [
|
||||
{"task_id": "seg_task_1", "audio_url": "", "duration": 0.0, "file_size": 0},
|
||||
{"task_id": "seg_task_2", "audio_url": "", "duration": 0.0, "file_size": 0},
|
||||
]
|
||||
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = _make_job()
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo)
|
||||
job = workflow.start_synthesis("test_job_seg")
|
||||
|
||||
assert job.status == TTSJobStatus.PROCESSING
|
||||
assert "segment_task_ids" in job.metadata
|
||||
assert job.metadata["segment_task_ids"] == ["seg_task_1", "seg_task_2"]
|
||||
|
||||
def test_segment_submit_failure_marks_failed(self) -> None:
|
||||
"""分段提交失败时标记 job 为 failed。"""
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
service.submit_synthesize_task.side_effect = CosyVoiceError("API error")
|
||||
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = _make_job()
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo)
|
||||
job = workflow.start_synthesis("test_job_seg")
|
||||
|
||||
assert job.status == TTSJobStatus.FAILED
|
||||
|
||||
|
||||
class TestUploadMergedToOSS:
|
||||
"""测试 _upload_merged_to_oss 辅助方法。"""
|
||||
|
||||
def test_success(self) -> None:
|
||||
"""成功上传返回 URL 和 key。"""
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/merged.mp3"
|
||||
|
||||
workflow = _make_workflow(storage=storage)
|
||||
url, key = workflow._upload_merged_to_oss(b"audio data", "user_001", "job_001", "mp3")
|
||||
|
||||
assert url == "https://oss.example.com/merged.mp3"
|
||||
assert key == "tts-outputs/user_001/job_001.mp3"
|
||||
storage.upload_file.assert_called_once()
|
||||
call_args = storage.upload_file.call_args
|
||||
assert call_args[1]["content_type"] == "audio/mpeg"
|
||||
|
||||
def test_failure_returns_empty(self) -> None:
|
||||
"""上传失败返回空字符串。"""
|
||||
storage = MagicMock()
|
||||
storage.upload_file.side_effect = Exception("OSS error")
|
||||
|
||||
workflow = _make_workflow(storage=storage)
|
||||
url, key = workflow._upload_merged_to_oss(b"audio data", "user_001", "job_001", "mp3")
|
||||
|
||||
assert url == ""
|
||||
assert key == ""
|
||||
|
||||
|
||||
class TestHandleSegmentFailure:
|
||||
"""测试 _handle_segment_failure。"""
|
||||
|
||||
def test_marks_job_failed(self) -> None:
|
||||
"""标记 job 为 failed 并更新。"""
|
||||
repo = MagicMock()
|
||||
job = _make_job(status=TTSJobStatus.PROCESSING)
|
||||
repo.get.return_value = job
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(repo=repo)
|
||||
workflow._handle_segment_failure(job, "分段 1 合成失败")
|
||||
|
||||
assert job.status == TTSJobStatus.FAILED
|
||||
assert "分段 1 合成失败" in job.error_message
|
||||
repo.update.assert_called_once()
|
||||
|
||||
|
||||
class TestPollSegmentTasks:
|
||||
"""测试 _poll_segment_tasks 异步轮询。"""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.time")
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_all_segments_done(self, mock_httpx: MagicMock, mock_time: MagicMock) -> None:
|
||||
"""所有分段完成后合并并标记完成。"""
|
||||
# Mock time.monotonic 让循环只执行一次
|
||||
mock_time.monotonic.side_effect = [0.0, 1.0, 2.0]
|
||||
mock_time.sleep = MagicMock()
|
||||
|
||||
# Mock 下载分段音频
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"seg audio"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
service.poll_synthesize_task.side_effect = [
|
||||
{"audio_url": "https://temp.com/seg1.mp3", "duration": 2.0, "file_size": 100},
|
||||
{"audio_url": "https://temp.com/seg2.mp3", "duration": 3.0, "file_size": 200},
|
||||
]
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/merged.mp3"
|
||||
|
||||
repo = MagicMock()
|
||||
job = _make_job(
|
||||
status=TTSJobStatus.PROCESSING,
|
||||
metadata={
|
||||
"segment_task_ids": ["task_1", "task_2"],
|
||||
"segment_audio_urls": ["", ""],
|
||||
"segment_format": "mp3",
|
||||
},
|
||||
)
|
||||
repo.get.return_value = job
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo, storage=storage)
|
||||
|
||||
with patch("packages.application.tts_job.workflow.AudioMerger") as MockMerger:
|
||||
mock_merger = MagicMock()
|
||||
mock_merger.merge.return_value = b"merged data"
|
||||
MockMerger.return_value = mock_merger
|
||||
|
||||
result = workflow._poll_segment_tasks(job)
|
||||
|
||||
assert result.status == TTSJobStatus.COMPLETED
|
||||
|
||||
@patch("packages.application.tts_job.workflow.time")
|
||||
def test_segment_poll_failure(self, mock_time: MagicMock) -> None:
|
||||
"""分段轮询失败时标记 job failed。"""
|
||||
mock_time.monotonic.side_effect = [0.0, 1.0]
|
||||
mock_time.sleep = MagicMock()
|
||||
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
service.poll_synthesize_task.side_effect = CosyVoiceError("Poll failed")
|
||||
|
||||
repo = MagicMock()
|
||||
job = _make_job(
|
||||
status=TTSJobStatus.PROCESSING,
|
||||
metadata={
|
||||
"segment_task_ids": ["task_1"],
|
||||
"segment_audio_urls": [""],
|
||||
},
|
||||
)
|
||||
repo.get.return_value = job
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo)
|
||||
result = workflow._poll_segment_tasks(job)
|
||||
|
||||
assert result.status == TTSJobStatus.FAILED
|
||||
|
||||
|
||||
class TestPollAndProcessSynthesisSegmentDetection:
|
||||
"""测试 poll_and_process_synthesis 正确识别分段任务。"""
|
||||
|
||||
def test_detects_segment_task(self) -> None:
|
||||
"""metadata 中有 segment_task_ids 时走分段轮询路径。"""
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
|
||||
repo = MagicMock()
|
||||
job = _make_job(
|
||||
status=TTSJobStatus.PROCESSING,
|
||||
metadata={
|
||||
"segment_task_ids": ["task_1", "task_2"],
|
||||
"segment_audio_urls": ["", ""],
|
||||
},
|
||||
)
|
||||
repo.get.return_value = job
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo)
|
||||
|
||||
with patch.object(workflow, "_poll_segment_tasks") as mock_poll:
|
||||
mock_poll.return_value = job
|
||||
workflow.poll_and_process_synthesis("test_job_seg")
|
||||
mock_poll.assert_called_once()
|
||||
|
||||
def test_normal_task_no_segment(self) -> None:
|
||||
"""普通任务不走分段路径。"""
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
service.poll_synthesize_task.return_value = {
|
||||
"audio_url": "https://temp.com/audio.mp3",
|
||||
"duration": 5.0,
|
||||
"file_size": 5000,
|
||||
}
|
||||
|
||||
repo = MagicMock()
|
||||
job = _make_job(
|
||||
status=TTSJobStatus.PROCESSING,
|
||||
metadata={"cosyvoice_task_id": "task_normal"},
|
||||
)
|
||||
repo.get.return_value = job
|
||||
repo.update.side_effect = lambda j: j
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/audio.mp3"
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo, storage=storage)
|
||||
|
||||
with patch("packages.application.tts_job.workflow.httpx") as mock_httpx:
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"audio"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
|
||||
result = workflow.poll_and_process_synthesis("test_job_seg")
|
||||
|
||||
assert result.status == TTSJobStatus.COMPLETED
|
||||
@@ -0,0 +1,239 @@
|
||||
"""P2 WebSocket 流式合成单元测试。
|
||||
|
||||
覆盖:
|
||||
- TTSStreamingService 流式合成逻辑
|
||||
- 短文本流式合成
|
||||
- 长文本分段流式合成
|
||||
- 错误处理
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
|
||||
from packages.application.tts_job.streaming_service import (
|
||||
TTSStreamingError,
|
||||
TTSStreamingService,
|
||||
)
|
||||
|
||||
|
||||
class MockWebSocket:
|
||||
"""Mock WebSocket for testing."""
|
||||
|
||||
def __init__(self):
|
||||
self.sent_json = []
|
||||
self.sent_bytes = []
|
||||
self.accepted = False
|
||||
|
||||
async def accept(self):
|
||||
self.accepted = True
|
||||
|
||||
async def send_json(self, data):
|
||||
self.sent_json.append(data)
|
||||
|
||||
async def send_bytes(self, data):
|
||||
self.sent_bytes.append(data)
|
||||
|
||||
async def receive_json(self):
|
||||
return {}
|
||||
|
||||
|
||||
class TestTTSStreamingService:
|
||||
"""测试 TTS 流式合成服务。"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_text_returns_error(self):
|
||||
"""空文本返回错误。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
ws = MockWebSocket()
|
||||
params = {"text": "", "voice_id": "test_voice"}
|
||||
|
||||
await service.synthesize_and_stream(ws, params)
|
||||
|
||||
assert len(ws.sent_json) == 1
|
||||
assert ws.sent_json[0]["type"] == "error"
|
||||
assert "文本不能为空" in ws.sent_json[0]["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_too_long_returns_error(self):
|
||||
"""超长文本返回错误。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
ws = MockWebSocket()
|
||||
params = {"text": "x" * 10001, "voice_id": "test_voice"}
|
||||
|
||||
await service.synthesize_and_stream(ws, params)
|
||||
|
||||
assert len(ws.sent_json) == 1
|
||||
assert ws.sent_json[0]["type"] == "error"
|
||||
assert "文本过长" in ws.sent_json[0]["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_short_text_stream_success(self):
|
||||
"""短文本流式合成成功。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
cosyvoice.submit_synthesize_task.return_value = {
|
||||
"task_id": "task_1",
|
||||
"audio_url": "https://temp.com/audio.mp3",
|
||||
"duration": 5.0,
|
||||
"file_size": 10000,
|
||||
}
|
||||
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
ws = MockWebSocket()
|
||||
params = {"text": "测试文本", "voice_id": "test_voice", "format": "mp3"}
|
||||
|
||||
with patch.object(service, "_download_audio", return_value=b"fake audio data"):
|
||||
await service.synthesize_and_stream(ws, params)
|
||||
|
||||
# 验证发送了 started 帧
|
||||
assert ws.sent_json[0]["type"] == "started"
|
||||
assert ws.sent_json[0]["segment_count"] == 1
|
||||
|
||||
# 验证发送了二进制音频数据
|
||||
assert len(ws.sent_bytes) > 0
|
||||
|
||||
# 验证发送了 done 帧
|
||||
assert ws.sent_json[-1]["type"] == "done"
|
||||
assert ws.sent_json[-1]["format"] == "mp3"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_short_text_cosyvoice_error(self):
|
||||
"""短文本合成时 CosyVoice 报错。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
cosyvoice.submit_synthesize_task.side_effect = CosyVoiceError("API error")
|
||||
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
ws = MockWebSocket()
|
||||
params = {"text": "测试文本", "voice_id": "test_voice"}
|
||||
|
||||
await service.synthesize_and_stream(ws, params)
|
||||
|
||||
# 验证发送了 started 帧和 error 帧
|
||||
assert len(ws.sent_json) == 2
|
||||
assert ws.sent_json[0]["type"] == "started"
|
||||
assert ws.sent_json[1]["type"] == "error"
|
||||
assert "API error" in ws.sent_json[1]["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_short_text_no_audio_url(self):
|
||||
"""短文本合成未返回 audio_url。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
cosyvoice.submit_synthesize_task.return_value = {
|
||||
"task_id": "task_1",
|
||||
"audio_url": "",
|
||||
}
|
||||
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
ws = MockWebSocket()
|
||||
params = {"text": "测试文本", "voice_id": "test_voice"}
|
||||
|
||||
await service.synthesize_and_stream(ws, params)
|
||||
|
||||
# 验证发送了 started 帧和 error 帧
|
||||
assert len(ws.sent_json) == 2
|
||||
assert ws.sent_json[0]["type"] == "started"
|
||||
assert ws.sent_json[1]["type"] == "error"
|
||||
assert "未返回音频 URL" in ws.sent_json[1]["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_long_text_stream_success(self):
|
||||
"""长文本分段流式合成成功。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
# 每个分段都返回 audio_url
|
||||
cosyvoice.submit_synthesize_task.side_effect = [
|
||||
{"task_id": "", "audio_url": "https://temp.com/seg1.mp3", "duration": 2.0},
|
||||
{"task_id": "", "audio_url": "https://temp.com/seg2.mp3", "duration": 3.0},
|
||||
]
|
||||
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
ws = MockWebSocket()
|
||||
# 超过 500 字的文本
|
||||
params = {"text": "x" * 600, "voice_id": "test_voice", "format": "mp3"}
|
||||
|
||||
with patch.object(service, "_download_audio", return_value=b"segment audio"):
|
||||
await service.synthesize_and_stream(ws, params)
|
||||
|
||||
# 验证发送了 started 帧(多段)
|
||||
assert ws.sent_json[0]["type"] == "started"
|
||||
assert ws.sent_json[0]["segment_count"] >= 2
|
||||
|
||||
# 验证发送了 segment_done 帧
|
||||
segment_done_count = sum(1 for msg in ws.sent_json if msg["type"] == "segment_done")
|
||||
assert segment_done_count >= 2
|
||||
|
||||
# 验证发送了 done 帧
|
||||
assert ws.sent_json[-1]["type"] == "done"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_long_text_segment_failure(self):
|
||||
"""长文本分段合成失败。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
cosyvoice.submit_synthesize_task.side_effect = CosyVoiceError("Segment error")
|
||||
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
ws = MockWebSocket()
|
||||
params = {"text": "x" * 600, "voice_id": "test_voice"}
|
||||
|
||||
await service.synthesize_and_stream(ws, params)
|
||||
|
||||
# 验证发送了错误帧
|
||||
error_msgs = [msg for msg in ws.sent_json if msg["type"] == "error"]
|
||||
assert len(error_msgs) > 0
|
||||
assert "合成失败" in error_msgs[0]["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_audio_chunks(self):
|
||||
"""音频分块推送。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
ws = MockWebSocket()
|
||||
audio_data = b"x" * 10000 # 10KB
|
||||
|
||||
total = await service._stream_audio_chunks(ws, audio_data)
|
||||
|
||||
assert total == 10000
|
||||
# 验证分块发送(4KB per chunk)
|
||||
assert len(ws.sent_bytes) == 3 # 4096 + 4096 + 1808
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_json_suppresses_exceptions(self):
|
||||
"""_send_json 抑制异常。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
ws = MockWebSocket()
|
||||
ws.send_json = AsyncMock(side_effect=Exception("Send failed"))
|
||||
|
||||
# 不应该抛出异常
|
||||
await service._send_json(ws, {"type": "test"})
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_audio(self):
|
||||
"""下载音频数据。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
with patch("packages.application.tts_job.streaming_service.httpx") as mock_httpx:
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"audio data"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
|
||||
result = service._download_audio("https://example.com/audio.mp3")
|
||||
|
||||
assert result == b"audio data"
|
||||
mock_httpx.get.assert_called_once()
|
||||
Reference in New Issue
Block a user