Merge branch 'main' of https://gitee.com/wg123/video-gen
This commit is contained in:
@@ -0,0 +1,31 @@
|
||||
"""private portrait validate once
|
||||
|
||||
Revision ID: e7038d02c355
|
||||
Revises: 6fc75582f6f9
|
||||
Create Date: 2026-07-07 10:47:00.436752
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = 'e7038d02c355'
|
||||
down_revision: Union[str, None] = '6fc75582f6f9'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_index('uq_private_portrait_asset_groups_one_active_project', 'private_portrait_asset_groups', ['project_id'], unique=True, postgresql_where=sa.text("deleted_at IS NULL AND status = 'active'"))
|
||||
op.create_index('uq_private_portrait_validate_sessions_one_group_active_project', 'private_portrait_validate_sessions', ['project_id'], unique=True, postgresql_where=sa.text("status = 'group_active'"))
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_index('uq_private_portrait_validate_sessions_one_group_active_project', table_name='private_portrait_validate_sessions', postgresql_where=sa.text("status = 'group_active'"))
|
||||
op.drop_index('uq_private_portrait_asset_groups_one_active_project', table_name='private_portrait_asset_groups', postgresql_where=sa.text("deleted_at IS NULL AND status = 'active'"))
|
||||
# ### end Alembic commands ###
|
||||
@@ -49,6 +49,7 @@ from app.services.admin_credit_record_service import list_admin_credit_records
|
||||
from app.services.notification import create_notification
|
||||
from app.services.auth import hash_password, verify_password
|
||||
from app.services.operation_log import log_operation
|
||||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls
|
||||
from app.services.resource_signed_url_service import build_resource_signed_url
|
||||
from app.services.payment import sync_pending_orders, process_refund
|
||||
from app.services.resource_capacity_service import batch_get_user_resource_capacity_usage, get_user_resource_capacity_usage
|
||||
@@ -1449,15 +1450,15 @@ async def admin_list_generation_records(
|
||||
query = query.offset(offset).limit(page_size)
|
||||
result = await db.execute(query)
|
||||
rows = result.all()
|
||||
refs_map = await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{record.id: json.loads(record.media_references) if record.media_references else None for record, _username, _project_name, _industry, _industry_label in rows},
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
items = []
|
||||
for record, username, project_name, industry, industry_label in rows:
|
||||
refs = None
|
||||
if record.media_references:
|
||||
try:
|
||||
refs = json.loads(record.media_references)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
refs = None
|
||||
refs = refs_map.get(record.id)
|
||||
items.append({
|
||||
"id": record.id,
|
||||
"user_id": record.user_id,
|
||||
|
||||
@@ -33,6 +33,7 @@ from app.services.resource_accounting_service import (
|
||||
record_generation_record_generated_resource,
|
||||
safe_file_size,
|
||||
)
|
||||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls
|
||||
from app.services.resource_signed_url_service import build_resource_signed_url
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.services.generation_billing_service import (
|
||||
@@ -59,9 +60,9 @@ router = APIRouter(prefix="/generation-records", tags=["generation"])
|
||||
logger = logging.getLogger("videogen")
|
||||
|
||||
|
||||
def _record_to_out(record: GenerationRecord, project_name: str) -> GenerationRecordOut:
|
||||
refs = None
|
||||
if record.media_references:
|
||||
def _record_to_out(record: GenerationRecord, project_name: str, refs_override: list[dict] | None = None) -> GenerationRecordOut:
|
||||
refs = refs_override
|
||||
if refs is None and record.media_references:
|
||||
try:
|
||||
refs = json.loads(record.media_references)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
@@ -191,13 +192,18 @@ async def list_records(
|
||||
|
||||
result = await db.execute(query)
|
||||
rows = result.all()
|
||||
refs_map = await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{record.id: json.loads(record.media_references) if record.media_references else None for record, _project_name in rows},
|
||||
user_id=current_user.id,
|
||||
)
|
||||
|
||||
return {
|
||||
"total": int(total),
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"items": [
|
||||
_record_to_out(record, project_name)
|
||||
_record_to_out(record, project_name, refs_override=refs_map.get(record.id))
|
||||
for record, project_name in rows
|
||||
],
|
||||
}
|
||||
@@ -240,11 +246,12 @@ async def optimize(
|
||||
row = existing.first()
|
||||
if row:
|
||||
record, project_name = row
|
||||
refs = await resolve_private_portrait_reference_display_urls(db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id)
|
||||
return OptimizeResult(
|
||||
optimized_prompt=record.optimized_prompt or "",
|
||||
text_credits_cost=record.text_credits_cost or 0.00,
|
||||
text_tokens_used=record.text_tokens_used or 0,
|
||||
record=_record_to_out(record, project_name),
|
||||
record=_record_to_out(record, project_name, refs_override=refs),
|
||||
)
|
||||
|
||||
# Check project exists and belongs to user
|
||||
@@ -370,11 +377,12 @@ async def optimize(
|
||||
record.text_tokens_used = token_usage["total_tokens"]
|
||||
await db.flush()
|
||||
|
||||
refs = await resolve_private_portrait_reference_display_urls(db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id)
|
||||
return OptimizeResult(
|
||||
optimized_prompt=optimized,
|
||||
text_credits_cost=round(text_credits, 2),
|
||||
# text_tokens_used=token_usage["total_tokens"],
|
||||
record=_record_to_out(record, project.name),
|
||||
record=_record_to_out(record, project.name, refs_override=refs),
|
||||
)
|
||||
|
||||
|
||||
@@ -502,7 +510,8 @@ async def generate(
|
||||
)
|
||||
await db.flush()
|
||||
|
||||
return _record_to_out(record, project_name)
|
||||
refs = await resolve_private_portrait_reference_display_urls(db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id)
|
||||
return _record_to_out(record, project_name, refs_override=refs)
|
||||
|
||||
|
||||
@router.post("/{record_id}/retry")
|
||||
@@ -578,7 +587,8 @@ async def retry_generation(
|
||||
)
|
||||
await db.flush()
|
||||
|
||||
return _record_to_out(record, project_name)
|
||||
refs = await resolve_private_portrait_reference_display_urls(db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id)
|
||||
return _record_to_out(record, project_name, refs_override=refs)
|
||||
|
||||
|
||||
@router.put("/{record_id}/prompt")
|
||||
|
||||
@@ -36,6 +36,7 @@ from app.services.generation_billing_service import (
|
||||
from app.services.generation_history_delete_service import batch_delete_generation_history_items
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.tasks.celery_app import celery_app
|
||||
|
||||
@@ -173,7 +174,8 @@ async def create_task(
|
||||
await db.commit()
|
||||
raise HTTPException(status_code=503, detail="任务队列投递失败,请稍后重试")
|
||||
|
||||
return record_to_out(task)
|
||||
refs = await resolve_private_portrait_reference_display_urls(db, record_to_out(task).media_references, user_id=current_user.id)
|
||||
return record_to_out(task, media_references=refs)
|
||||
|
||||
|
||||
@router.get(
|
||||
@@ -280,7 +282,12 @@ async def list_tasks(
|
||||
)
|
||||
else:
|
||||
items_sorted = items
|
||||
return GenerationAITaskListOut(total=total, items=[record_to_out(task=i, is_admin=is_admin) for i in items_sorted])
|
||||
refs_map = await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{item.id: record_to_out(task=item, is_admin=is_admin).media_references for item in items_sorted},
|
||||
user_id=None if is_admin else current_user.id,
|
||||
)
|
||||
return GenerationAITaskListOut(total=total, items=[record_to_out(task=i, is_admin=is_admin, media_references=refs_map.get(i.id)) for i in items_sorted])
|
||||
|
||||
|
||||
@router.get(
|
||||
@@ -521,7 +528,8 @@ async def get_task(
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
return record_to_out(task)
|
||||
refs = await resolve_private_portrait_reference_display_urls(db, record_to_out(task).media_references, user_id=current_user.id)
|
||||
return record_to_out(task, media_references=refs)
|
||||
|
||||
|
||||
@router.delete(
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from urllib.parse import unquote
|
||||
from urllib.parse import urlencode, unquote
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.responses import RedirectResponse
|
||||
@@ -12,6 +12,7 @@ from app.enums.private_portrait import (
|
||||
PrivatePortraitEventSource,
|
||||
PrivatePortraitEventStatus,
|
||||
PrivatePortraitEventType,
|
||||
PrivatePortraitProjectStatus,
|
||||
PrivatePortraitRemoteDeleteStatus,
|
||||
)
|
||||
from app.models.private_portrait import PrivatePortraitAsset, PrivatePortraitProject
|
||||
@@ -22,6 +23,7 @@ from app.schemas.private_portrait import (
|
||||
PrivatePortraitDeleteOut,
|
||||
PrivatePortraitConfigOut,
|
||||
PrivatePortraitProjectCreate,
|
||||
PrivatePortraitProjectCreateWithValidateOut,
|
||||
PrivatePortraitProjectListOut,
|
||||
PrivatePortraitProjectOut,
|
||||
PrivatePortraitProjectUpdate,
|
||||
@@ -88,10 +90,20 @@ async def get_my_private_portrait_config(current_user: User = Depends(get_curren
|
||||
return await get_user_private_portrait_config(db, user_id=current_user.id)
|
||||
|
||||
|
||||
@router.post("/private-portrait/projects", response_model=PrivatePortraitProjectOut)
|
||||
@router.post("/private-portrait/projects", response_model=PrivatePortraitProjectCreateWithValidateOut)
|
||||
async def create_private_portrait_project(payload: PrivatePortraitProjectCreate, current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
|
||||
project = await create_project(db, user_id=current_user.id, payload=payload)
|
||||
out = project_to_out(project)
|
||||
session = await create_validate_session(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
project_id=project.id,
|
||||
callback_redirect_url=payload.callback_redirect_url,
|
||||
)
|
||||
out = PrivatePortraitProjectCreateWithValidateOut(
|
||||
project=project_to_out(project),
|
||||
validate_session=validate_session_to_out(session),
|
||||
poll_interval_ms=2000,
|
||||
)
|
||||
await db.commit()
|
||||
return out
|
||||
|
||||
@@ -105,10 +117,11 @@ async def list_private_portrait_projects(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
items, total = await list_projects(db, user_id=current_user.id, page=page, page_size=page_size, keyword=keyword, status=status)
|
||||
query_status = status or PrivatePortraitProjectStatus.ACTIVE.value
|
||||
items, total = await list_projects(db, user_id=current_user.id, page=page, page_size=page_size, keyword=keyword, status=query_status)
|
||||
await refresh_project_counters(db, [item.id for item in items])
|
||||
await db.commit()
|
||||
items, total = await list_projects(db, user_id=current_user.id, page=page, page_size=page_size, keyword=keyword, status=status)
|
||||
items, total = await list_projects(db, user_id=current_user.id, page=page, page_size=page_size, keyword=keyword, status=query_status)
|
||||
return PrivatePortraitProjectListOut(items=[project_to_out(item) for item in items], total=total, page=page, page_size=page_size)
|
||||
|
||||
|
||||
@@ -159,15 +172,19 @@ async def private_portrait_validate_callback(session_id: str, request: Request,
|
||||
params.pop("session_id", None)
|
||||
params.pop("redirect_url", None)
|
||||
session = await handle_validate_callback(db, session_id=session_id, query_params=params)
|
||||
redirect_session_id = session.id
|
||||
redirect_status = session.status
|
||||
redirect_result_code = session.result_code or ""
|
||||
redirect_params = {
|
||||
"session_id": session.id,
|
||||
"status": session.status,
|
||||
"resultCode": session.result_code or "",
|
||||
}
|
||||
if session.remote_group_id:
|
||||
redirect_params["remote_group_id"] = session.remote_group_id
|
||||
response = {"session_id": session.id, "status": session.status, "resultCode": session.result_code, "remote_group_id": session.remote_group_id}
|
||||
await db.commit()
|
||||
if redirect_url:
|
||||
sep = "&" if "?" in redirect_url else "?"
|
||||
url = f"{unquote(redirect_url)}{sep}session_id={redirect_session_id}&status={redirect_status}&resultCode={redirect_result_code}"
|
||||
return RedirectResponse(url=url)
|
||||
base_url = unquote(redirect_url)
|
||||
sep = "&" if "?" in base_url else "?"
|
||||
return RedirectResponse(url=f"{base_url}{sep}{urlencode(redirect_params)}")
|
||||
return response
|
||||
|
||||
|
||||
|
||||
@@ -56,7 +56,9 @@ class ArkPrivatePortraitAction(str, Enum):
|
||||
|
||||
|
||||
class PrivatePortraitProjectStatus(str, Enum):
|
||||
VALIDATING = "validating"
|
||||
ACTIVE = "active"
|
||||
VALIDATE_FAILED = "validate_failed"
|
||||
DELETED = "deleted"
|
||||
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, ForeignKey, Index, String, Text
|
||||
from sqlalchemy import DateTime, ForeignKey, Index, String, Text, text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.enums.private_portrait import (
|
||||
@@ -23,6 +23,12 @@ class PrivatePortraitAssetGroup(Base, TimestampMixin, SoftDeleteMixin):
|
||||
Index("idx_private_portrait_asset_groups_project_status", "project_id", "status"),
|
||||
Index("idx_private_portrait_asset_groups_remote_delete_status", "remote_delete_status"),
|
||||
Index("idx_private_portrait_asset_groups_remote_project_name", "remote_project_name"),
|
||||
Index(
|
||||
"uq_private_portrait_asset_groups_one_active_project",
|
||||
"project_id",
|
||||
unique=True,
|
||||
postgresql_where=text("deleted_at IS NULL AND status = 'active'"),
|
||||
),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, ForeignKey, Index, String, Text
|
||||
from sqlalchemy import DateTime, ForeignKey, Index, String, Text, text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.enums.private_portrait import PrivatePortraitValidateSessionStatus
|
||||
@@ -17,6 +17,12 @@ class PrivatePortraitValidateSession(Base, TimestampMixin):
|
||||
Index("idx_private_portrait_validate_sessions_user_project", "user_id", "project_id"),
|
||||
Index("idx_private_portrait_validate_sessions_byted_token", "byted_token"),
|
||||
Index("idx_private_portrait_validate_sessions_status_created", "status", "created_at"),
|
||||
Index(
|
||||
"uq_private_portrait_validate_sessions_one_group_active_project",
|
||||
"project_id",
|
||||
unique=True,
|
||||
postgresql_where=text("status = 'group_active'"),
|
||||
),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
|
||||
@@ -56,6 +56,21 @@ class GenerationAIReference(BaseModel):
|
||||
description="后端回填的火山 Asset ID。前端传入时不可信,创建任务时以后端查库为准",
|
||||
examples=["asset-20260318071009-xxxxx"],
|
||||
)
|
||||
provider_url: str | None = Field(
|
||||
None,
|
||||
description="供应商专用素材地址。真人素材生成时通常为 asset://remote_asset_id;仅用于后端排查和生成链路,不用于前端预览",
|
||||
examples=["asset://asset-20260318071009-xxxxx"],
|
||||
)
|
||||
display_url: str | None = Field(
|
||||
None,
|
||||
description="前端展示用素材地址。真人素材历史响应会回填为可预览图片/视频/音频 URL",
|
||||
examples=["/uploads/images/2026/07/06/demo.jpg"],
|
||||
)
|
||||
preview_url: str | None = Field(
|
||||
None,
|
||||
description="前端预览用素材地址;通常与 display_url 一致",
|
||||
examples=["/uploads/images/2026/07/06/demo.jpg"],
|
||||
)
|
||||
|
||||
|
||||
class GenerationAITaskCreate(BaseModel):
|
||||
|
||||
@@ -22,6 +22,7 @@ class PrivatePortraitAdminConfigUpdate(BaseModel):
|
||||
class PrivatePortraitProjectCreate(BaseModel):
|
||||
name: str = Field(..., min_length=1, max_length=128)
|
||||
description: str | None = Field(None, max_length=2000)
|
||||
callback_redirect_url: str | None = Field(None, description="项目创建时真人认证完成后的手机端提示页地址")
|
||||
|
||||
|
||||
class PrivatePortraitProjectUpdate(BaseModel):
|
||||
@@ -80,6 +81,14 @@ class PrivatePortraitValidateSessionOut(BaseModel):
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
|
||||
|
||||
class PrivatePortraitProjectCreateWithValidateOut(BaseModel):
|
||||
project: PrivatePortraitProjectOut
|
||||
validate_session: PrivatePortraitValidateSessionOut
|
||||
poll_interval_ms: int = Field(default=2000, description="PC 端轮询认证状态的建议间隔,单位毫秒")
|
||||
|
||||
|
||||
class PrivatePortraitAssetGroupOut(BaseModel):
|
||||
id: str
|
||||
user_id: str | None = None
|
||||
|
||||
@@ -54,7 +54,7 @@ from app.services.generation_history_meta_service import (
|
||||
build_empty_history_meta,
|
||||
)
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
|
||||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls, resolve_private_portrait_references
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
IMAGE_DEFAULT_SIZE = "2K"
|
||||
@@ -88,6 +88,22 @@ def _parse_json(text: str | None):
|
||||
return None
|
||||
|
||||
|
||||
async def _resolve_task_reference_display_map(db: AsyncSession, tasks: list[ChatGenerationTask], *, user_id: str | None = None) -> dict[str, list[dict] | None]:
|
||||
return await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{task.id: _parse_json(task.media_references) for task in tasks},
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
|
||||
async def _resolve_generation_record_reference_display_map(db: AsyncSession, records: list[GenerationRecord], *, user_id: str | None = None) -> dict[str, list[dict] | None]:
|
||||
return await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{record.id: _parse_json(record.media_references) for record in records},
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
|
||||
async def _get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine:
|
||||
query = select(ImageEngine).where(ImageEngine.is_active == True)
|
||||
if engine_id:
|
||||
@@ -462,8 +478,9 @@ def record_to_out(
|
||||
generated_resource_id: str | None = None,
|
||||
file_name: str | None = None,
|
||||
history_meta: GenerationHistoryMeta | None = None,
|
||||
media_references: list[dict] | None = None,
|
||||
) -> GenerationAITaskOut:
|
||||
refs = _parse_json(task.media_references)
|
||||
refs = media_references if media_references is not None else _parse_json(task.media_references)
|
||||
snapshot = engine_snapshot_out(_parse_json(task.engine_snapshot_json))
|
||||
|
||||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||||
@@ -688,8 +705,9 @@ def generation_record_to_history_out(
|
||||
project_name: str | None = None,
|
||||
generated_resource_id: str | None = None,
|
||||
file_name: str | None = None,
|
||||
media_references: list[dict] | None = None,
|
||||
) -> GenerationAIRecordHistoryItemOut:
|
||||
refs = _parse_json(record.media_references)
|
||||
refs = media_references if media_references is not None else _parse_json(record.media_references)
|
||||
return GenerationAIRecordHistoryItemOut(
|
||||
id=record.id,
|
||||
source_type="generation_record",
|
||||
@@ -814,6 +832,8 @@ async def list_generation_record_history_grouped_days(
|
||||
source_ids=all_record_ids,
|
||||
resource_type=gen_type,
|
||||
)
|
||||
all_records = [record for _generated_day, _day_total, rows in raw_groups for record, _project_name in rows]
|
||||
reference_display_map = await _resolve_generation_record_reference_display_map(db, all_records, user_id=user_id)
|
||||
|
||||
groups = [
|
||||
{
|
||||
@@ -825,6 +845,7 @@ async def list_generation_record_history_grouped_days(
|
||||
project_name,
|
||||
generated_resource_id=resource_info_map.get(record.id, {}).get("resource_id"),
|
||||
file_name=resource_info_map.get(record.id, {}).get("file_name"),
|
||||
media_references=reference_display_map.get(record.id),
|
||||
)
|
||||
for record, project_name in rows
|
||||
],
|
||||
@@ -887,6 +908,7 @@ async def list_generation_record_history_day_items(
|
||||
source_ids=[record.id for record, _project_name in rows],
|
||||
resource_type=gen_type,
|
||||
)
|
||||
reference_display_map = await _resolve_generation_record_reference_display_map(db, [record for record, _project_name in rows], user_id=user_id)
|
||||
|
||||
return {
|
||||
"generated_date": target_day.strftime("%Y-%m-%d"),
|
||||
@@ -899,6 +921,7 @@ async def list_generation_record_history_day_items(
|
||||
project_name,
|
||||
generated_resource_id=resource_info_map.get(record.id, {}).get("resource_id"),
|
||||
file_name=resource_info_map.get(record.id, {}).get("file_name"),
|
||||
media_references=reference_display_map.get(record.id),
|
||||
)
|
||||
for record, project_name in rows
|
||||
],
|
||||
@@ -989,6 +1012,8 @@ async def list_generation_history_grouped_days(
|
||||
source=source,
|
||||
chat_task_ids=all_task_ids,
|
||||
)
|
||||
all_tasks = [task for _generated_day, _day_total, tasks in raw_groups for task in tasks]
|
||||
reference_display_map = await _resolve_task_reference_display_map(db, all_tasks, user_id=user_id)
|
||||
|
||||
groups = [
|
||||
{
|
||||
@@ -1000,6 +1025,7 @@ async def list_generation_history_grouped_days(
|
||||
generated_resource_id=resource_info_map.get(task.id, {}).get("resource_id"),
|
||||
file_name=resource_info_map.get(task.id, {}).get("file_name"),
|
||||
history_meta=history_meta_map.get(task.id),
|
||||
media_references=reference_display_map.get(task.id),
|
||||
)
|
||||
for task in tasks
|
||||
],
|
||||
@@ -1082,6 +1108,7 @@ async def list_generation_history_day_items(
|
||||
source=source,
|
||||
chat_task_ids=task_ids,
|
||||
)
|
||||
reference_display_map = await _resolve_task_reference_display_map(db, tasks, user_id=user_id)
|
||||
|
||||
return {
|
||||
"generated_date": target_day.strftime("%Y-%m-%d"),
|
||||
@@ -1094,6 +1121,7 @@ async def list_generation_history_day_items(
|
||||
generated_resource_id=resource_info_map.get(task.id, {}).get("resource_id"),
|
||||
file_name=resource_info_map.get(task.id, {}).get("file_name"),
|
||||
history_meta=history_meta_map.get(task.id),
|
||||
media_references=reference_display_map.get(task.id),
|
||||
)
|
||||
for task in tasks
|
||||
],
|
||||
|
||||
@@ -23,6 +23,7 @@ from app.enums.private_portrait import (
|
||||
PrivatePortraitEventSource,
|
||||
PrivatePortraitEventStatus,
|
||||
PrivatePortraitEventType,
|
||||
PrivatePortraitProjectStatus,
|
||||
PrivatePortraitRemoteDeleteStatus,
|
||||
PrivatePortraitValidateSessionStatus,
|
||||
)
|
||||
@@ -165,8 +166,69 @@ def asset_to_out(asset: PrivatePortraitAsset, *, project_name: str | None = None
|
||||
)
|
||||
|
||||
|
||||
async def _get_existing_active_group(db: AsyncSession, *, project_id: str) -> PrivatePortraitAssetGroup | None:
|
||||
return (
|
||||
await db.execute(
|
||||
select(PrivatePortraitAssetGroup)
|
||||
.where(
|
||||
PrivatePortraitAssetGroup.project_id == project_id,
|
||||
PrivatePortraitAssetGroup.status == PrivatePortraitAssetGroupStatus.ACTIVE.value,
|
||||
PrivatePortraitAssetGroup.deleted_at.is_(None),
|
||||
)
|
||||
.order_by(PrivatePortraitAssetGroup.created_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
|
||||
|
||||
async def _ensure_project_can_validate(db: AsyncSession, *, project: PrivatePortraitProject) -> PrivatePortraitValidateSession | None:
|
||||
active_group = await _get_existing_active_group(db, project_id=project.id)
|
||||
if active_group or project.status == PrivatePortraitProjectStatus.ACTIVE.value:
|
||||
raise HTTPException(status_code=409, detail="该真人素材项目已完成认证,不能重复认证")
|
||||
|
||||
success_session = (
|
||||
await db.execute(
|
||||
select(PrivatePortraitValidateSession)
|
||||
.where(
|
||||
PrivatePortraitValidateSession.project_id == project.id,
|
||||
PrivatePortraitValidateSession.status == PrivatePortraitValidateSessionStatus.GROUP_ACTIVE.value,
|
||||
)
|
||||
.order_by(PrivatePortraitValidateSession.updated_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
if success_session:
|
||||
raise HTTPException(status_code=409, detail="该真人素材项目已完成认证,不能重复认证")
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
pending_session = (
|
||||
await db.execute(
|
||||
select(PrivatePortraitValidateSession)
|
||||
.where(
|
||||
PrivatePortraitValidateSession.project_id == project.id,
|
||||
PrivatePortraitValidateSession.status.in_(
|
||||
[
|
||||
PrivatePortraitValidateSessionStatus.CREATED.value,
|
||||
PrivatePortraitValidateSessionStatus.CALLBACK_SUCCESS.value,
|
||||
]
|
||||
),
|
||||
PrivatePortraitValidateSession.expired_at.is_not(None),
|
||||
PrivatePortraitValidateSession.expired_at > now,
|
||||
)
|
||||
.order_by(PrivatePortraitValidateSession.created_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
).scalar_one_or_none()
|
||||
return pending_session
|
||||
|
||||
|
||||
async def create_validate_session(db: AsyncSession, *, user_id: str, project_id: str, callback_redirect_url: str | None = None) -> PrivatePortraitValidateSession:
|
||||
project = await get_user_project(db, user_id=user_id, project_id=project_id)
|
||||
reusable_session = await _ensure_project_can_validate(db, project=project)
|
||||
if reusable_session:
|
||||
return reusable_session
|
||||
|
||||
project.status = PrivatePortraitProjectStatus.VALIDATING.value
|
||||
session = PrivatePortraitValidateSession(
|
||||
id=generate_id(),
|
||||
user_id=user_id,
|
||||
@@ -192,11 +254,11 @@ async def create_validate_session(db: AsyncSession, *, user_id: str, project_id:
|
||||
except Exception as exc:
|
||||
session.status = PrivatePortraitValidateSessionStatus.FAILED.value
|
||||
session.error_message = _exception_message(exc)
|
||||
project.status = PrivatePortraitProjectStatus.VALIDATE_FAILED.value
|
||||
await db.flush()
|
||||
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_SESSION_CREATE_FAILED.value, source=PrivatePortraitEventSource.API.value, user_id=user_id, project_id=project.id, session_id=session.id, exc=exc)
|
||||
raise
|
||||
|
||||
|
||||
async def get_validate_session(db: AsyncSession, *, user_id: str | None, session_id: str) -> PrivatePortraitValidateSession:
|
||||
filters = [PrivatePortraitValidateSession.id == session_id]
|
||||
if user_id is not None:
|
||||
@@ -209,6 +271,9 @@ async def get_validate_session(db: AsyncSession, *, user_id: str | None, session
|
||||
|
||||
async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_params: dict[str, Any]) -> PrivatePortraitValidateSession:
|
||||
session = await get_validate_session(db, user_id=None, session_id=session_id)
|
||||
if session.status == PrivatePortraitValidateSessionStatus.GROUP_ACTIVE.value:
|
||||
return session
|
||||
|
||||
session.raw_callback_json = _json(query_params)
|
||||
session.result_code = str(query_params.get("resultCode") or query_params.get("result_code") or "") or None
|
||||
session.algorithm_base_resp_code = str(query_params.get("algorithmBaseRespCode") or query_params.get("algorithm_base_resp_code") or "") or None
|
||||
@@ -218,9 +283,13 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
|
||||
session.byted_token = str(token)
|
||||
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_CALLBACK_RECEIVED.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.CALLBACK.value, user_id=session.user_id, project_id=session.project_id, session_id=session.id, detail={"query_params": query_params, "remote_project_name": session.remote_project_name})
|
||||
|
||||
project = (await db.execute(select(PrivatePortraitProject).where(PrivatePortraitProject.id == session.project_id).limit(1))).scalar_one_or_none()
|
||||
|
||||
if session.result_code != PRIVATE_PORTRAIT_SUCCESS_RESULT_CODE:
|
||||
session.status = PrivatePortraitValidateSessionStatus.CALLBACK_FAILED.value
|
||||
session.error_message = f"真人认证失败:resultCode={session.result_code}"
|
||||
if project and project.status != PrivatePortraitProjectStatus.ACTIVE.value:
|
||||
project.status = PrivatePortraitProjectStatus.VALIDATE_FAILED.value
|
||||
await db.flush()
|
||||
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_CALLBACK_FAILED.value, event_status=PrivatePortraitEventStatus.FAILED.value, source=PrivatePortraitEventSource.CALLBACK.value, user_id=session.user_id, project_id=session.project_id, session_id=session.id, error=session.error_message)
|
||||
return session
|
||||
@@ -229,10 +298,22 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
|
||||
if not session.byted_token:
|
||||
session.status = PrivatePortraitValidateSessionStatus.FAILED.value
|
||||
session.error_message = "Callback 未返回 BytedToken"
|
||||
if project and project.status != PrivatePortraitProjectStatus.ACTIVE.value:
|
||||
project.status = PrivatePortraitProjectStatus.VALIDATE_FAILED.value
|
||||
await db.flush()
|
||||
raise HTTPException(status_code=400, detail=session.error_message)
|
||||
|
||||
try:
|
||||
existing_group = await _get_existing_active_group(db, project_id=session.project_id)
|
||||
if existing_group:
|
||||
session.remote_group_id = existing_group.remote_group_id
|
||||
session.status = PrivatePortraitValidateSessionStatus.GROUP_ACTIVE.value
|
||||
if project:
|
||||
project.status = PrivatePortraitProjectStatus.ACTIVE.value
|
||||
await db.flush()
|
||||
await db.refresh(session)
|
||||
return session
|
||||
|
||||
log_operation_event(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_GET_RESULT_START.value, event_status=PrivatePortraitEventStatus.PENDING.value, source=PrivatePortraitEventSource.CALLBACK.value, user_id=session.user_id, project_id=session.project_id, session_id=session.id, detail={"remote_project_name": session.remote_project_name})
|
||||
resp = await ArkPrivateAssetClient().get_visual_validate_result(project_name=session.remote_project_name, byted_token=session.byted_token)
|
||||
group_id = resp.get("GroupId") or resp.get("groupId")
|
||||
@@ -242,7 +323,8 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
|
||||
session.status = PrivatePortraitValidateSessionStatus.GROUP_ACTIVE.value
|
||||
session.raw_response_json = _json(resp)
|
||||
|
||||
project = (await db.execute(select(PrivatePortraitProject).where(PrivatePortraitProject.id == session.project_id).limit(1))).scalar_one()
|
||||
if not project:
|
||||
raise RuntimeError("真人素材项目不存在")
|
||||
remote_group_name = _remote_group_name(session.user_id, project.name)
|
||||
group = PrivatePortraitAssetGroup(
|
||||
id=generate_id(),
|
||||
@@ -256,6 +338,7 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
|
||||
raw_response_json=_json(resp),
|
||||
)
|
||||
db.add(group)
|
||||
project.status = PrivatePortraitProjectStatus.ACTIVE.value
|
||||
await db.flush()
|
||||
try:
|
||||
await ArkPrivateAssetClient().update_asset_group(project_name=session.remote_project_name, group_id=group_id, name=remote_group_name, title=remote_group_name, description=project.description)
|
||||
@@ -269,12 +352,16 @@ async def handle_validate_callback(db: AsyncSession, *, session_id: str, query_p
|
||||
except Exception as exc:
|
||||
session.status = PrivatePortraitValidateSessionStatus.FAILED.value
|
||||
session.error_message = _exception_message(exc)
|
||||
if project and project.status != PrivatePortraitProjectStatus.ACTIVE.value:
|
||||
project.status = PrivatePortraitProjectStatus.VALIDATE_FAILED.value
|
||||
await db.flush()
|
||||
log_operation_error(domain=DOMAIN, event_type=PrivatePortraitEventType.VALIDATE_GET_RESULT_FAILED.value, source=PrivatePortraitEventSource.CALLBACK.value, user_id=session.user_id, project_id=session.project_id, session_id=session.id, exc=exc)
|
||||
raise
|
||||
|
||||
|
||||
async def get_project_active_group(db: AsyncSession, *, user_id: str, project_id: str) -> PrivatePortraitAssetGroup:
|
||||
project = await get_user_project(db, user_id=user_id, project_id=project_id)
|
||||
if project.status != PrivatePortraitProjectStatus.ACTIVE.value:
|
||||
raise HTTPException(status_code=400, detail="请先完成真人授权认证,再上传素材")
|
||||
result = await db.execute(select(PrivatePortraitAssetGroup).where(PrivatePortraitAssetGroup.user_id == user_id, PrivatePortraitAssetGroup.project_id == project_id, PrivatePortraitAssetGroup.status == PrivatePortraitAssetGroupStatus.ACTIVE.value, PrivatePortraitAssetGroup.deleted_at.is_(None)).order_by(PrivatePortraitAssetGroup.created_at.desc()).limit(1))
|
||||
group = result.scalar_one_or_none()
|
||||
if not group:
|
||||
@@ -284,6 +371,8 @@ async def get_project_active_group(db: AsyncSession, *, user_id: str, project_id
|
||||
|
||||
async def create_asset(db: AsyncSession, *, user_id: str, project_id: str, payload: PrivatePortraitAssetCreate) -> PrivatePortraitAsset:
|
||||
project = await get_user_project(db, user_id=user_id, project_id=project_id)
|
||||
if project.status != PrivatePortraitProjectStatus.ACTIVE.value:
|
||||
raise HTTPException(status_code=400, detail="项目正在真人认证或认证未通过,不能上传素材")
|
||||
if payload.asset_type != PrivatePortraitAssetType.IMAGE.value:
|
||||
raise HTTPException(status_code=400, detail="第一版真人素材库仅开放 Image 图片素材")
|
||||
user = await _lock_user_for_upload(db, user_id=user_id)
|
||||
|
||||
@@ -17,7 +17,7 @@ from app.enums.private_portrait import (
|
||||
PrivatePortraitProjectStatus,
|
||||
PrivatePortraitRemoteDeleteStatus,
|
||||
)
|
||||
from app.models.private_portrait import PrivatePortraitAsset, PrivatePortraitAssetGroup, PrivatePortraitProject, PrivatePortraitValidateSession
|
||||
from app.models.private_portrait import PrivatePortraitAsset, PrivatePortraitAssetGroup, PrivatePortraitProject
|
||||
from app.schemas.private_portrait import PrivatePortraitProjectCreate, PrivatePortraitProjectOut, PrivatePortraitProjectUpdate
|
||||
from app.services.operation_log_service import log_operation_event
|
||||
from app.utils.id_gen import generate_id
|
||||
@@ -74,19 +74,19 @@ async def create_project(db: AsyncSession, *, user_id: str, payload: PrivatePort
|
||||
name_slug=slug,
|
||||
remote_project_name=PRIVATE_PORTRAIT_REMOTE_PROJECT_NAME,
|
||||
description=payload.description,
|
||||
status=PrivatePortraitProjectStatus.ACTIVE.value,
|
||||
status=PrivatePortraitProjectStatus.VALIDATING.value,
|
||||
)
|
||||
db.add(project)
|
||||
await db.flush()
|
||||
log_operation_event(
|
||||
domain=DOMAIN,
|
||||
event_type=PrivatePortraitEventType.PROJECT_CREATE.value,
|
||||
event_status=PrivatePortraitEventStatus.SUCCESS.value,
|
||||
event_status=PrivatePortraitEventStatus.PENDING.value,
|
||||
source=PrivatePortraitEventSource.API.value,
|
||||
user_id=user_id,
|
||||
project_id=project.id,
|
||||
message="创建真人素材项目",
|
||||
detail={"name": project.name, "remote_project_name": project.remote_project_name},
|
||||
message="创建待认证真人素材项目",
|
||||
detail={"name": project.name, "remote_project_name": project.remote_project_name, "status": project.status},
|
||||
)
|
||||
return project
|
||||
|
||||
@@ -109,7 +109,12 @@ async def update_project(db: AsyncSession, *, user_id: str, project_id: str, pay
|
||||
if payload.description is not None:
|
||||
project.description = payload.description
|
||||
if payload.status is not None:
|
||||
if payload.status not in {PrivatePortraitProjectStatus.ACTIVE.value}:
|
||||
allowed_statuses = {
|
||||
PrivatePortraitProjectStatus.VALIDATING.value,
|
||||
PrivatePortraitProjectStatus.ACTIVE.value,
|
||||
PrivatePortraitProjectStatus.VALIDATE_FAILED.value,
|
||||
}
|
||||
if payload.status not in allowed_statuses:
|
||||
raise HTTPException(status_code=400, detail="项目状态不支持")
|
||||
project.status = payload.status
|
||||
await db.flush()
|
||||
@@ -133,7 +138,15 @@ async def update_project(db: AsyncSession, *, user_id: str, project_id: str, pay
|
||||
return project
|
||||
|
||||
|
||||
async def list_projects(db: AsyncSession, *, user_id: str | None, page: int = 1, page_size: int = 20, keyword: str | None = None, status: str | None = None) -> tuple[list[PrivatePortraitProject], int]:
|
||||
async def list_projects(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
user_id: str | None,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
keyword: str | None = None,
|
||||
status: str | None = None,
|
||||
) -> tuple[list[PrivatePortraitProject], int]:
|
||||
page = max(1, page)
|
||||
page_size = min(max(1, page_size), 100)
|
||||
filters = [PrivatePortraitProject.deleted_at.is_(None)]
|
||||
|
||||
@@ -4,7 +4,7 @@ from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy import or_, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.private_portrait import (
|
||||
@@ -58,6 +58,162 @@ def _normalize_ref_type(value: Any) -> str:
|
||||
return str(value or "").strip().lower()
|
||||
|
||||
|
||||
def _remote_asset_id_from_asset_uri(url: Any) -> str | None:
|
||||
value = str(url or "").strip()
|
||||
if not value.startswith(PRIVATE_PORTRAIT_ASSET_URI_PREFIX):
|
||||
return None
|
||||
remote_asset_id = value[len(PRIVATE_PORTRAIT_ASSET_URI_PREFIX):].strip()
|
||||
return remote_asset_id or None
|
||||
|
||||
|
||||
def _asset_display_url(asset: PrivatePortraitAsset) -> str | None:
|
||||
# preview_url 是本地上传预览,remote_url 是火山 GetAsset 返回的远程资源 URL;source_url 是兜底公网上传地址。
|
||||
return asset.preview_url or asset.remote_url or asset.source_url or None
|
||||
|
||||
|
||||
def _fill_private_portrait_reference_display_fields(ref: Any, asset: PrivatePortraitAsset) -> None:
|
||||
provider_url = str(_ref_get(ref, "provider_url") or _ref_get(ref, "url") or "").strip()
|
||||
if provider_url.startswith(PRIVATE_PORTRAIT_ASSET_URI_PREFIX):
|
||||
_ref_set(ref, "provider_url", provider_url)
|
||||
|
||||
display_url = _asset_display_url(asset)
|
||||
if display_url:
|
||||
# 返回给前端的 url 必须可预览;供应商专用 asset:// 保留到 provider_url,避免管理后台和客户端展示黑图。
|
||||
_ref_set(ref, "url", display_url)
|
||||
_ref_set(ref, "display_url", display_url)
|
||||
_ref_set(ref, "preview_url", display_url)
|
||||
|
||||
_ref_set(ref, "source", PrivatePortraitReferenceSource.PRIVATE_PORTRAIT_ASSET.value)
|
||||
_ref_set(ref, "private_asset_id", asset.id)
|
||||
if asset.remote_asset_id:
|
||||
_ref_set(ref, "remote_asset_id", asset.remote_asset_id)
|
||||
expected_ref_type = _ASSET_TYPE_TO_REFERENCE_TYPE.get(asset.asset_type)
|
||||
if expected_ref_type:
|
||||
_ref_set(ref, "type", expected_ref_type)
|
||||
if not _ref_get(ref, "name") and asset.name:
|
||||
_ref_set(ref, "name", asset.name)
|
||||
|
||||
|
||||
async def resolve_private_portrait_reference_display_urls(
|
||||
db: AsyncSession,
|
||||
media_references: list[Any] | None,
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
) -> list[Any] | None:
|
||||
"""把历史响应里的 asset:// 引用补成前端可预览 URL。
|
||||
|
||||
生成任务入库时 url 使用 asset://remote_asset_id 传给供应商;但客户端/管理后台展示不能直接用
|
||||
asset://。这里批量根据 private_asset_id 或 asset://remote_asset_id 查本地素材,并把响应中的 url
|
||||
改成 preview_url/remote_url/source_url,同时保留 provider_url=asset://... 供排查。
|
||||
"""
|
||||
if not media_references:
|
||||
return media_references
|
||||
|
||||
refs = deepcopy(media_references)
|
||||
private_asset_ids: list[str] = []
|
||||
remote_asset_ids: list[str] = []
|
||||
for ref in refs:
|
||||
source = _ref_get(ref, "source")
|
||||
private_asset_id = _ref_get(ref, "private_asset_id")
|
||||
remote_asset_id = _ref_get(ref, "remote_asset_id") or _remote_asset_id_from_asset_uri(_ref_get(ref, "url"))
|
||||
if source == PrivatePortraitReferenceSource.PRIVATE_PORTRAIT_ASSET.value or remote_asset_id:
|
||||
if private_asset_id:
|
||||
private_asset_ids.append(str(private_asset_id))
|
||||
if remote_asset_id:
|
||||
remote_asset_ids.append(str(remote_asset_id))
|
||||
|
||||
private_asset_ids = list(dict.fromkeys(private_asset_ids))
|
||||
remote_asset_ids = list(dict.fromkeys(remote_asset_ids))
|
||||
if not private_asset_ids and not remote_asset_ids:
|
||||
return refs
|
||||
|
||||
filters = []
|
||||
if private_asset_ids:
|
||||
filters.append(PrivatePortraitAsset.id.in_(private_asset_ids))
|
||||
if remote_asset_ids:
|
||||
filters.append(PrivatePortraitAsset.remote_asset_id.in_(remote_asset_ids))
|
||||
|
||||
stmt = select(PrivatePortraitAsset).where(or_(*filters))
|
||||
if user_id is not None:
|
||||
stmt = stmt.where(PrivatePortraitAsset.user_id == user_id)
|
||||
rows = await db.execute(stmt)
|
||||
assets = list(rows.scalars().all())
|
||||
by_id = {asset.id: asset for asset in assets}
|
||||
by_remote_id = {asset.remote_asset_id: asset for asset in assets if asset.remote_asset_id}
|
||||
|
||||
for ref in refs:
|
||||
private_asset_id = str(_ref_get(ref, "private_asset_id") or "").strip()
|
||||
remote_asset_id = str(_ref_get(ref, "remote_asset_id") or _remote_asset_id_from_asset_uri(_ref_get(ref, "url")) or "").strip()
|
||||
asset = by_id.get(private_asset_id) or by_remote_id.get(remote_asset_id)
|
||||
if not asset:
|
||||
continue
|
||||
_fill_private_portrait_reference_display_fields(ref, asset)
|
||||
|
||||
return refs
|
||||
|
||||
|
||||
async def batch_resolve_private_portrait_reference_display_urls(
|
||||
db: AsyncSession,
|
||||
references_by_key: dict[Any, list[Any] | None],
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
) -> dict[Any, list[Any] | None]:
|
||||
if not references_by_key:
|
||||
return {}
|
||||
|
||||
copied: dict[Any, list[Any] | None] = {
|
||||
key: deepcopy(refs) if refs else refs
|
||||
for key, refs in references_by_key.items()
|
||||
}
|
||||
private_asset_ids: list[str] = []
|
||||
remote_asset_ids: list[str] = []
|
||||
|
||||
for refs in copied.values():
|
||||
if not refs:
|
||||
continue
|
||||
for ref in refs:
|
||||
source = _ref_get(ref, "source")
|
||||
private_asset_id = _ref_get(ref, "private_asset_id")
|
||||
remote_asset_id = _ref_get(ref, "remote_asset_id") or _remote_asset_id_from_asset_uri(_ref_get(ref, "url"))
|
||||
if source == PrivatePortraitReferenceSource.PRIVATE_PORTRAIT_ASSET.value or remote_asset_id:
|
||||
if private_asset_id:
|
||||
private_asset_ids.append(str(private_asset_id))
|
||||
if remote_asset_id:
|
||||
remote_asset_ids.append(str(remote_asset_id))
|
||||
|
||||
private_asset_ids = list(dict.fromkeys(private_asset_ids))
|
||||
remote_asset_ids = list(dict.fromkeys(remote_asset_ids))
|
||||
if not private_asset_ids and not remote_asset_ids:
|
||||
return copied
|
||||
|
||||
filters = []
|
||||
if private_asset_ids:
|
||||
filters.append(PrivatePortraitAsset.id.in_(private_asset_ids))
|
||||
if remote_asset_ids:
|
||||
filters.append(PrivatePortraitAsset.remote_asset_id.in_(remote_asset_ids))
|
||||
|
||||
stmt = select(PrivatePortraitAsset).where(or_(*filters))
|
||||
if user_id is not None:
|
||||
stmt = stmt.where(PrivatePortraitAsset.user_id == user_id)
|
||||
rows = await db.execute(stmt)
|
||||
assets = list(rows.scalars().all())
|
||||
by_id = {asset.id: asset for asset in assets}
|
||||
by_remote_id = {asset.remote_asset_id: asset for asset in assets if asset.remote_asset_id}
|
||||
|
||||
for refs in copied.values():
|
||||
if not refs:
|
||||
continue
|
||||
for ref in refs:
|
||||
private_asset_id = str(_ref_get(ref, "private_asset_id") or "").strip()
|
||||
remote_asset_id = str(_ref_get(ref, "remote_asset_id") or _remote_asset_id_from_asset_uri(_ref_get(ref, "url")) or "").strip()
|
||||
asset = by_id.get(private_asset_id) or by_remote_id.get(remote_asset_id)
|
||||
if not asset:
|
||||
continue
|
||||
_fill_private_portrait_reference_display_fields(ref, asset)
|
||||
|
||||
return copied
|
||||
|
||||
|
||||
async def resolve_private_portrait_references(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
@@ -123,10 +279,16 @@ async def resolve_private_portrait_references(
|
||||
if expected_ref_type and ref_type and ref_type != expected_ref_type:
|
||||
raise HTTPException(status_code=400, detail=f"真人素材类型不匹配:引用为 {ref_type},素材为 {expected_ref_type}")
|
||||
|
||||
provider_url = f"{PRIVATE_PORTRAIT_ASSET_URI_PREFIX}{asset.remote_asset_id}"
|
||||
_ref_set(ref, "source", PrivatePortraitReferenceSource.PRIVATE_PORTRAIT_ASSET.value)
|
||||
_ref_set(ref, "private_asset_id", asset.id)
|
||||
_ref_set(ref, "remote_asset_id", asset.remote_asset_id)
|
||||
_ref_set(ref, "url", f"{PRIVATE_PORTRAIT_ASSET_URI_PREFIX}{asset.remote_asset_id}")
|
||||
_ref_set(ref, "url", provider_url)
|
||||
_ref_set(ref, "provider_url", provider_url)
|
||||
display_url = _asset_display_url(asset)
|
||||
if display_url:
|
||||
_ref_set(ref, "display_url", display_url)
|
||||
_ref_set(ref, "preview_url", display_url)
|
||||
if expected_ref_type:
|
||||
_ref_set(ref, "type", expected_ref_type)
|
||||
if not _ref_get(ref, "name") and asset.name:
|
||||
|
||||
+713
File diff suppressed because one or more lines are too long
+209
-1
File diff suppressed because one or more lines are too long
Vendored
+1
@@ -33,5 +33,6 @@
|
||||
</head>
|
||||
<body>
|
||||
<div id="root"></div>
|
||||
|
||||
</body>
|
||||
</html>
|
||||
|
||||
@@ -34,6 +34,7 @@ import PopularPage from './pages/PopularPage';
|
||||
import CreativePlazaPage from './pages/CreativePlazaPage';
|
||||
import TeamManagementPage from './pages/TeamManagementPage';
|
||||
import JoinTeamPage from './pages/JoinTeamPage';
|
||||
import PrivatePortraitAuthorizeResult from './pages/PrivatePortraitAuthorizeResult';
|
||||
import { useAuthStore } from './store/useAuthStore';
|
||||
const ProtectedRoute = ({ children }: { children: React.ReactNode }) => {
|
||||
const { user, loading, checkAuth } = useAuthStore();
|
||||
|
||||
@@ -8,7 +8,7 @@ import type {
|
||||
User, CreditRecord, Project, GenerationRecord, OptimizeParams, GenerateParams, OptimizeResult,
|
||||
Industry, IndustryConfig, AdminUser, AdminStats, ModelConfig, SystemConfig, AdminNotification,
|
||||
PrivatePortraitConfig, PrivatePortraitProjectListOut, PrivatePortraitProject, PrivatePortraitValidateSession,
|
||||
PrivatePortraitAssetListOut, PrivatePortraitAsset, PrivatePortraitSelectableAssetListOut,
|
||||
PrivatePortraitProjectCreateWithValidateOut, PrivatePortraitAssetListOut, PrivatePortraitAsset, PrivatePortraitSelectableAssetListOut,
|
||||
} from '../types';
|
||||
const USE_MOCK = import.meta.env.VITE_USE_MOCK === 'true';
|
||||
// ── Auth ──────────────────────────────────────────────────
|
||||
@@ -756,8 +756,12 @@ export async function getPrivatePortraitProjects(params: { page?: number; pageSi
|
||||
return api.get<PrivatePortraitProjectListOut>(`/private-portrait/projects?${query.toString()}`);
|
||||
}
|
||||
|
||||
export async function createPrivatePortraitProject(payload: { name: string; description?: string | null }): Promise<PrivatePortraitProject> {
|
||||
return api.post<PrivatePortraitProject>('/private-portrait/projects', payload);
|
||||
export async function createPrivatePortraitProject(payload: { name: string; description?: string | null; callbackRedirectUrl?: string | null }): Promise<PrivatePortraitProjectCreateWithValidateOut> {
|
||||
return api.post<PrivatePortraitProjectCreateWithValidateOut>('/private-portrait/projects', {
|
||||
name: payload.name,
|
||||
description: payload.description || null,
|
||||
callback_redirect_url: payload.callbackRedirectUrl || null,
|
||||
});
|
||||
}
|
||||
|
||||
export async function updatePrivatePortraitProject(projectId: string, payload: { name?: string; description?: string | null; status?: string }): Promise<PrivatePortraitProject> {
|
||||
|
||||
@@ -1,22 +1,47 @@
|
||||
import React, { useEffect, useState } from 'react';
|
||||
import { Button, Card, Col, Form, Input, Modal, Row, Space, Typography, message } from 'antd';
|
||||
import { PlusOutlined, ReloadOutlined } from '@ant-design/icons';
|
||||
import type { PrivatePortraitProject } from '../../../types';
|
||||
import { createPrivatePortraitProject, getPrivatePortraitProjects } from '../../../api';
|
||||
import React, { useEffect, useRef, useState } from 'react';
|
||||
import { Button, Card, Col, Form, Input, Modal, QRCode, Row, Space, Spin, Typography, message } from 'antd';
|
||||
import { CheckCircleOutlined, PlusOutlined, ReloadOutlined } from '@ant-design/icons';
|
||||
import type { PrivatePortraitProject, PrivatePortraitValidateSession } from '../../../types';
|
||||
import { createPrivatePortraitProject, getPrivatePortraitProjects, getPrivatePortraitValidateSession } from '../../../api';
|
||||
import PrivatePortraitProjectList from './ProjectList';
|
||||
import PrivatePortraitProjectDetail from './ProjectDetail';
|
||||
|
||||
const VALIDATE_SUCCESS_STATUS = 'group_active';
|
||||
const POLL_INTERVAL_FALLBACK = 2000;
|
||||
|
||||
const PrivatePortraitLibraryPanel: React.FC = () => {
|
||||
const [projects, setProjects] = useState<PrivatePortraitProject[]>([]);
|
||||
const [selected, setSelected] = useState<PrivatePortraitProject | null>(null);
|
||||
const [loading, setLoading] = useState(false);
|
||||
const [createOpen, setCreateOpen] = useState(false);
|
||||
const [creating, setCreating] = useState(false);
|
||||
const [createdProject, setCreatedProject] = useState<PrivatePortraitProject | null>(null);
|
||||
const [validateSession, setValidateSession] = useState<PrivatePortraitValidateSession | null>(null);
|
||||
const [polling, setPolling] = useState(false);
|
||||
const [form] = Form.useForm();
|
||||
const timerRef = useRef<number | null>(null);
|
||||
|
||||
const clearPollTimer = () => {
|
||||
if (timerRef.current) {
|
||||
window.clearInterval(timerRef.current);
|
||||
timerRef.current = null;
|
||||
}
|
||||
};
|
||||
|
||||
const resetCreateModal = () => {
|
||||
clearPollTimer();
|
||||
setCreateOpen(false);
|
||||
setCreating(false);
|
||||
setPolling(false);
|
||||
setCreatedProject(null);
|
||||
setValidateSession(null);
|
||||
form.resetFields();
|
||||
};
|
||||
|
||||
const loadProjects = async () => {
|
||||
setLoading(true);
|
||||
try {
|
||||
const res = await getPrivatePortraitProjects({ pageSize: 100 });
|
||||
const res = await getPrivatePortraitProjects({ pageSize: 100, status: 'active' });
|
||||
setProjects(res.items);
|
||||
setSelected((prev) => prev ? (res.items.find((item) => item.id === prev.id) || res.items[0] || null) : (res.items[0] || null));
|
||||
} catch (e: any) {
|
||||
@@ -27,27 +52,75 @@ const PrivatePortraitLibraryPanel: React.FC = () => {
|
||||
};
|
||||
|
||||
useEffect(() => { loadProjects(); }, []);
|
||||
useEffect(() => () => clearPollTimer(), []);
|
||||
|
||||
const finishCreateSuccess = async (projectId: string) => {
|
||||
clearPollTimer();
|
||||
setPolling(false);
|
||||
message.success('真人认证完成,项目组已创建成功');
|
||||
setCreateOpen(false);
|
||||
setValidateSession(null);
|
||||
setCreatedProject(null);
|
||||
form.resetFields();
|
||||
const res = await getPrivatePortraitProjects({ pageSize: 100, status: 'active' });
|
||||
setProjects(res.items);
|
||||
setSelected(res.items.find((item) => item.id === projectId) || res.items[0] || null);
|
||||
};
|
||||
|
||||
const startPolling = (sessionId: string, projectId: string, intervalMs: number) => {
|
||||
clearPollTimer();
|
||||
setPolling(true);
|
||||
const run = async () => {
|
||||
try {
|
||||
const next = await getPrivatePortraitValidateSession(sessionId);
|
||||
setValidateSession(next);
|
||||
if (next.status === VALIDATE_SUCCESS_STATUS) {
|
||||
await finishCreateSuccess(projectId);
|
||||
return;
|
||||
}
|
||||
if (['callback_failed', 'failed', 'expired'].includes(next.status)) {
|
||||
clearPollTimer();
|
||||
setPolling(false);
|
||||
message.error(next.errorMessage || '真人认证未完成,请重新创建项目组');
|
||||
}
|
||||
} catch (e: any) {
|
||||
clearPollTimer();
|
||||
setPolling(false);
|
||||
message.error(e?.message || '轮询真人认证状态失败');
|
||||
}
|
||||
};
|
||||
timerRef.current = window.setInterval(run, Math.max(1000, intervalMs || POLL_INTERVAL_FALLBACK));
|
||||
void run();
|
||||
};
|
||||
|
||||
const handleCreate = async () => {
|
||||
const values = await form.validateFields();
|
||||
setCreating(true);
|
||||
try {
|
||||
const project = await createPrivatePortraitProject(values);
|
||||
message.success('项目组已创建');
|
||||
setCreateOpen(false);
|
||||
form.resetFields();
|
||||
await loadProjects();
|
||||
setSelected(project);
|
||||
const callbackRedirectUrl = `${window.location.origin}/private-portrait-authorized`;
|
||||
const res = await createPrivatePortraitProject({ ...values, callbackRedirectUrl });
|
||||
setCreatedProject(res.project);
|
||||
setValidateSession(res.validateSession);
|
||||
message.success('请使用手机扫码完成人脸认证');
|
||||
if (res.validateSession?.id) {
|
||||
startPolling(res.validateSession.id, res.project.id, res.pollIntervalMs || POLL_INTERVAL_FALLBACK);
|
||||
}
|
||||
} catch (e: any) {
|
||||
message.error(e?.message || '创建项目组失败');
|
||||
message.error(e?.message || '创建项目组认证二维码失败');
|
||||
} finally {
|
||||
setCreating(false);
|
||||
}
|
||||
};
|
||||
|
||||
const h5Link = validateSession?.h5Link || '';
|
||||
const isSuccess = validateSession?.status === VALIDATE_SUCCESS_STATUS;
|
||||
|
||||
return (
|
||||
<div style={{ padding: 16 }}>
|
||||
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center', marginBottom: 16 }}>
|
||||
<div>
|
||||
<Typography.Title level={4} style={{ margin: 0 }}>真人素材库</Typography.Title>
|
||||
<Typography.Text type="secondary">管理真人授权项目组和已入库 Active 素材,AI 创作添加参考内容时可直接选择。</Typography.Text>
|
||||
<Typography.Text type="secondary">创建项目组时先完成真人认证,认证成功后项目组才会正式创建并可上传素材。</Typography.Text>
|
||||
</div>
|
||||
<Space>
|
||||
<Button icon={<ReloadOutlined />} onClick={loadProjects} loading={loading}>刷新</Button>
|
||||
@@ -68,15 +141,46 @@ const PrivatePortraitLibraryPanel: React.FC = () => {
|
||||
)}
|
||||
</Col>
|
||||
</Row>
|
||||
<Modal title="新建真人素材项目组" open={createOpen} onCancel={() => setCreateOpen(false)} onOk={handleCreate} okText="创建">
|
||||
<Form form={form} layout="vertical">
|
||||
<Form.Item label="项目组名称" name="name" rules={[{ required: true, message: '请输入项目组名称' }]}>
|
||||
<Input placeholder="例如:达人A、客户B、张三人像" />
|
||||
</Form.Item>
|
||||
<Form.Item label="描述" name="description">
|
||||
<Input.TextArea rows={3} placeholder="可选" />
|
||||
</Form.Item>
|
||||
</Form>
|
||||
<Modal
|
||||
title="新建真人素材项目组"
|
||||
open={createOpen}
|
||||
onCancel={resetCreateModal}
|
||||
footer={validateSession ? [<Button key="close" onClick={resetCreateModal}>关闭</Button>] : undefined}
|
||||
onOk={validateSession ? undefined : handleCreate}
|
||||
okText="开始认证并创建"
|
||||
confirmLoading={creating}
|
||||
maskClosable={!polling}
|
||||
>
|
||||
{!validateSession ? (
|
||||
<Form form={form} layout="vertical">
|
||||
<Form.Item label="项目组名称" name="name" rules={[{ required: true, message: '请输入项目组名称' }]}>
|
||||
<Input placeholder="例如:达人A、客户B、张三人像" />
|
||||
</Form.Item>
|
||||
<Form.Item label="描述" name="description">
|
||||
<Input.TextArea rows={3} placeholder="可选" />
|
||||
</Form.Item>
|
||||
<Typography.Paragraph type="secondary" style={{ marginBottom: 0 }}>
|
||||
点击后会生成真人认证二维码。手机扫码认证成功后,项目组才会出现在项目列表中。
|
||||
</Typography.Paragraph>
|
||||
</Form>
|
||||
) : (
|
||||
<Space direction="vertical" align="center" size={16} style={{ width: '100%' }}>
|
||||
{isSuccess ? (
|
||||
<CheckCircleOutlined style={{ fontSize: 54, color: '#22c55e' }} />
|
||||
) : h5Link ? (
|
||||
<QRCode value={h5Link} size={220} />
|
||||
) : (
|
||||
<Spin />
|
||||
)}
|
||||
<div style={{ textAlign: 'center' }}>
|
||||
<Typography.Title level={5} style={{ marginBottom: 8 }}>{createdProject?.name || '真人素材项目组'}</Typography.Title>
|
||||
<Typography.Text type={isSuccess ? 'success' : 'secondary'}>
|
||||
{isSuccess ? '认证成功,项目组正在刷新' : '请使用手机扫码完成人脸认证,成功后回到电脑端查看项目组。'}
|
||||
</Typography.Text>
|
||||
</div>
|
||||
{h5Link && !isSuccess && <Typography.Text copyable style={{ wordBreak: 'break-all' }}>{h5Link}</Typography.Text>}
|
||||
</Space>
|
||||
)}
|
||||
</Modal>
|
||||
</div>
|
||||
);
|
||||
|
||||
@@ -1,11 +1,10 @@
|
||||
import React, { useEffect, useState } from 'react';
|
||||
import { Button, Card, Popconfirm, Space, Typography, message } from 'antd';
|
||||
import { DeleteOutlined, ReloadOutlined, SafetyCertificateOutlined, UploadOutlined } from '@ant-design/icons';
|
||||
import { Button, Card, Popconfirm, Space, Tag, Typography, message } from 'antd';
|
||||
import { DeleteOutlined, ReloadOutlined, UploadOutlined } from '@ant-design/icons';
|
||||
import type { PrivatePortraitAsset, PrivatePortraitProject } from '../../../types';
|
||||
import { deletePrivatePortraitAsset, deletePrivatePortraitProject, getPrivatePortraitAssets, syncPrivatePortraitAsset } from '../../../api';
|
||||
import PrivatePortraitAssetGrid from './AssetGrid';
|
||||
import PrivatePortraitAssetUpload from './AssetUpload';
|
||||
import PrivatePortraitValidateModal from './ValidateModal';
|
||||
|
||||
interface Props {
|
||||
project: PrivatePortraitProject;
|
||||
@@ -17,7 +16,6 @@ const PrivatePortraitProjectDetail: React.FC<Props> = ({ project, onDeleted, onC
|
||||
const [assets, setAssets] = useState<PrivatePortraitAsset[]>([]);
|
||||
const [loading, setLoading] = useState(false);
|
||||
const [uploadOpen, setUploadOpen] = useState(false);
|
||||
const [validateOpen, setValidateOpen] = useState(false);
|
||||
|
||||
const loadAssets = async () => {
|
||||
setLoading(true);
|
||||
@@ -65,13 +63,14 @@ const PrivatePortraitProjectDetail: React.FC<Props> = ({ project, onDeleted, onC
|
||||
}
|
||||
};
|
||||
|
||||
const canUpload = project.status === 'active';
|
||||
|
||||
return (
|
||||
<Card
|
||||
title={<span>{project.name}</span>}
|
||||
title={<Space><span>{project.name}</span><Tag color={canUpload ? 'green' : 'processing'}>{project.status}</Tag></Space>}
|
||||
extra={(
|
||||
<Space>
|
||||
<Button icon={<SafetyCertificateOutlined />} onClick={() => setValidateOpen(true)}>真人授权</Button>
|
||||
<Button type="primary" icon={<UploadOutlined />} onClick={() => setUploadOpen(true)}>上传素材</Button>
|
||||
<Button type="primary" icon={<UploadOutlined />} disabled={!canUpload} onClick={() => setUploadOpen(true)}>上传素材</Button>
|
||||
<Button icon={<ReloadOutlined />} onClick={loadAssets} loading={loading}>刷新</Button>
|
||||
<Popconfirm title="确认删除这个真人素材项目组吗?" onConfirm={handleDeleteProject}>
|
||||
<Button danger icon={<DeleteOutlined />}>删除项目组</Button>
|
||||
@@ -81,9 +80,13 @@ const PrivatePortraitProjectDetail: React.FC<Props> = ({ project, onDeleted, onC
|
||||
style={{ borderRadius: 12 }}
|
||||
>
|
||||
<Typography.Paragraph style={{ color: '#64748b' }}>{project.description || '暂无描述'}</Typography.Paragraph>
|
||||
{!canUpload && (
|
||||
<Typography.Paragraph style={{ color: '#f97316' }}>
|
||||
项目组未完成真人认证,暂不能上传素材。请重新创建项目组并完成手机扫码认证。
|
||||
</Typography.Paragraph>
|
||||
)}
|
||||
<PrivatePortraitAssetGrid items={assets} loading={loading} onSync={handleSync} onDelete={handleDeleteAsset} />
|
||||
<PrivatePortraitAssetUpload projectId={project.id} open={uploadOpen} onClose={() => setUploadOpen(false)} onSuccess={() => { loadAssets(); onChanged(); }} />
|
||||
<PrivatePortraitValidateModal projectId={project.id} open={validateOpen} onClose={() => setValidateOpen(false)} onCreated={() => onChanged()} />
|
||||
</Card>
|
||||
);
|
||||
};
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
import React, { useMemo } from 'react';
|
||||
import { Button, Card, Result, Typography } from 'antd';
|
||||
|
||||
const successStatuses = new Set(['group_active', 'callback_success']);
|
||||
|
||||
const PrivatePortraitAuthorizeResult: React.FC = () => {
|
||||
const params = useMemo(() => new URLSearchParams(window.location.search), []);
|
||||
const status = params.get('status') || '';
|
||||
const resultCode = params.get('resultCode') || '';
|
||||
const isSuccess = successStatuses.has(status) || resultCode === '10000';
|
||||
|
||||
return (
|
||||
<div style={{ minHeight: '100vh', display: 'flex', alignItems: 'center', justifyContent: 'center', background: '#f8fafc', padding: 20 }}>
|
||||
<Card style={{ width: '100%', maxWidth: 520, borderRadius: 18 }}>
|
||||
<Result
|
||||
status={isSuccess ? 'success' : 'error'}
|
||||
title={isSuccess ? '真人认证已完成' : '真人认证未完成'}
|
||||
subTitle={isSuccess ? '请回到电脑端查看,项目组已创建成功。' : '请回到电脑端重新发起创建项目组。'}
|
||||
extra={[
|
||||
<Button key="close" type="primary" onClick={() => window.close()}>关闭页面</Button>,
|
||||
]}
|
||||
/>
|
||||
<Typography.Paragraph type="secondary" style={{ textAlign: 'center', marginBottom: 0 }}>
|
||||
当前状态:{status || '-'},结果码:{resultCode || '-'}
|
||||
</Typography.Paragraph>
|
||||
</Card>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
export default PrivatePortraitAuthorizeResult;
|
||||
@@ -135,6 +135,9 @@ export interface MediaReference {
|
||||
source?: string;
|
||||
private_asset_id?: string;
|
||||
remote_asset_id?: string;
|
||||
providerUrl?: string;
|
||||
displayUrl?: string;
|
||||
previewUrl?: string;
|
||||
}
|
||||
|
||||
export interface GenerationRecord {
|
||||
@@ -298,6 +301,13 @@ export interface PrivatePortraitValidateSession {
|
||||
updatedAt?: string | null;
|
||||
}
|
||||
|
||||
|
||||
export interface PrivatePortraitProjectCreateWithValidateOut {
|
||||
project: PrivatePortraitProject;
|
||||
validateSession: PrivatePortraitValidateSession;
|
||||
pollIntervalMs: number;
|
||||
}
|
||||
|
||||
export interface PrivatePortraitAsset {
|
||||
id: string;
|
||||
userId?: string | null;
|
||||
|
||||
Reference in New Issue
Block a user