merge main

This commit is contained in:
2026-08-11 10:16:38 +08:00
156 changed files with 22362 additions and 1211 deletions
+6
View File
@@ -12,6 +12,9 @@ from app.api.admin.llm_billing import router as llm_billing_router
from app.api.admin.menu_config import router as menu_config_router
from app.api.admin.upload import router as admin_upload_router
from app.api.admin.contact import router as admin_contact_router
from app.admin_api.api_keys import router as api_keys_admin_router
from app.admin_api.api_model_pricings import router as api_model_pricings_admin_router
from app.admin_api.vp_v3_quota import router as vp_v3_quota_admin_router
router = APIRouter()
router.include_router(video_prompt_schema_config_router)
@@ -26,3 +29,6 @@ router.include_router(llm_billing_router)
router.include_router(menu_config_router)
router.include_router(admin_upload_router)
router.include_router(admin_contact_router)
router.include_router(api_keys_admin_router)
router.include_router(api_model_pricings_admin_router)
router.include_router(vp_v3_quota_admin_router)
+4
View File
@@ -36,6 +36,8 @@ from app.api.v1.material_admin import router as material_admin_router
from app.api.v1.private_portrait import router as private_portrait_router
from app.api.v1.private_portrait_virtual import router as private_portrait_virtual_router
from app.api.v1.upload_resource import router as upload_resource_router
from app.api.v1.invoices import router as invoices_router
from app.api.v1.invoice_headers import router as invoice_headers_router
api_router = APIRouter()
api_router.include_router(auth_router)
@@ -74,3 +76,5 @@ api_router.include_router(material_admin_router)
api_router.include_router(private_portrait_router)
api_router.include_router(private_portrait_virtual_router)
api_router.include_router(upload_resource_router)
api_router.include_router(invoices_router)
api_router.include_router(invoice_headers_router)
+253 -23
View File
@@ -2,7 +2,7 @@ from datetime import datetime, timezone, timedelta
import json
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy import delete, func, or_, select, update
from sqlalchemy import and_, case, delete, func, or_, select, update
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_db, get_admin_user
@@ -63,6 +63,7 @@ from app.services.resource_signed_url_service import build_resource_signed_url
from app.services.payment import process_refund
from app.services.resource_capacity_service import batch_get_user_resource_capacity_usage, get_user_resource_capacity_usage
from app.services.team_service import batch_get_team_name_map, set_frontend_user_team
from app.schemas.invoice import InvoiceStatusUpdateRequest
from app.utils.id_gen import generate_id
@@ -1721,17 +1722,62 @@ async def update_system_config(
return config
@router.post("/system-configs/banner/reset", summary="重置活动横幅展示")
async def reset_banner(
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""递增 site_banner_version,使所有用户再次看到横幅。"""
from app.utils.id_gen import generate_id
result = await db.execute(select(SystemConfig).where(SystemConfig.key == "site_banner_version").limit(1))
config = result.scalar_one_or_none()
new_version = 1
if config:
try:
new_version = int(config.value or 0) + 1
except ValueError:
new_version = 1
config.value = str(new_version)
else:
config = SystemConfig(
id=generate_id(),
key="site_banner_version",
value=str(new_version),
description="活动横幅版本号,递增后所有用户重新看到横幅",
)
db.add(config)
await db.flush()
await log_operation(
db,
admin.id,
admin.username,
f"重置活动横幅 (版本 → {new_version})",
"POST",
"/admin/system-configs/banner/reset",
detail=json.dumps({"new_version": new_version}),
)
await db.commit()
await invalidate_system_config_cache(["site_banner_version"])
return {"site_banner_version": new_version}
# ── Operation Logs ──────────────────────────────────────
@router.get("/operation-logs")
async def list_operation_logs(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=500),
action: str | None = Query(None, description="按 action 过滤(前缀匹配)"),
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
query = select(OperationLog).order_by(OperationLog.created_at.desc())
count_query = select(func.count(OperationLog.id))
if action:
query = query.where(OperationLog.action.like(f"{action}%"))
count_query = count_query.where(OperationLog.action.like(f"{action}%"))
total = (await db.execute(count_query)).scalar() or 0
result = await db.execute(query.offset((page - 1) * page_size).limit(page_size))
items = result.scalars().all()
@@ -1774,17 +1820,26 @@ async def get_stats(
):
today_start = datetime.now(CST).replace(hour=0, minute=0, second=0, microsecond=0)
date_start: datetime
date_end: datetime
try:
if start_date:
date_start = datetime.strptime(start_date, "%Y-%m-%d").replace(tzinfo=CST)
else:
date_start = today_start
if end_date:
date_end = datetime.strptime(end_date, "%Y-%m-%d").replace(tzinfo=CST)
date_end = date_end.replace(hour=23, minute=59, second=59, microsecond=999999)
# 先构造完整的 naive 日期时刻,再一次性 attach tzinfo(避免分步 replace 丢 tzinfo
naive_end = datetime.strptime(end_date, "%Y-%m-%d").replace(
hour=23, minute=59, second=59, microsecond=999999,
)
date_end = naive_end.replace(tzinfo=CST)
else:
date_end = datetime.now(CST)
except:
# 合法性:end >= start
if date_end < date_start:
date_end = date_start.replace(hour=23, minute=59, second=59, microsecond=999999)
except (ValueError, TypeError):
# 只拦截日期解析错误,不吞掉 SQL/运行时异常(原裸 except 会吞所有错误导致用户看不到报错)
date_start = today_start
date_end = datetime.now(CST)
@@ -1827,20 +1882,45 @@ async def get_stats(
)
)).scalar() or 0
# 预扣占用不是实际消费;历史流水 charge_action 为空时仍按真实扣费兼容。
# 消费类(真实扣费 + 预扣占用):charge_action 为空时仍按真实扣费兼容hold 为预扣占用
credit_charge_action_filter = or_(
CreditRecord.charge_action.is_(None),
CreditRecord.charge_action == "charge",
CreditRecord.charge_action == "hold",
)
# 「仅真实扣费」filter 用于图表、模型使用次数等需要按实际产出(非预扣)统计的场景。
real_credit_charge_filter = or_(
CreditRecord.charge_action.is_(None),
CreditRecord.charge_action == "charge",
)
credits_consumed = (await db.execute(
select(func.coalesce(func.sum(func.abs(CreditRecord.amount)), 0)).where(
CreditRecord.type == "consume",
real_credit_charge_filter,
# 核心数据「消耗积分」= 净消耗 = 真实消费 + 预扣占用 - 真实退款 - 预扣释放。
# 说明:
# hold(预扣占用):type=consumecharge_action='hold'amount<0
# hold_release(预扣释放退回):type=refundcharge_action='hold_release'amount>0
# (账本 L256 强校验:hold_release.type 必须是 'refund',不是 consume
# charge(真实扣费):type=consumecharge_action='charge' 或 NULL(历史)amount<0
# refund(真实退款):type=refundcharge_action='refund' 或 NULL(历史兼容)amount>0
# 因此 type=refund 天然包含「真实退款 + 预扣释放退回」两类子流水。
_stats_real_and_hold = case(
(and_(CreditRecord.type == "consume", credit_charge_action_filter), func.abs(CreditRecord.amount)),
else_=0,
)
_stats_refund_and_release = case(
(CreditRecord.type == "refund", func.abs(CreditRecord.amount)),
else_=0,
)
_net_row = (await db.execute(
select(
func.coalesce(func.sum(_stats_real_and_hold), 0),
func.coalesce(func.sum(_stats_refund_and_release), 0),
).where(
CreditRecord.type.in_(["consume", "refund"]),
CreditRecord.created_at >= date_start,
CreditRecord.created_at <= date_end,
)
)).scalar() or 0
)).one()
credits_consumed = round(max(float(_net_row[0] or 0) - float(_net_row[1] or 0), 0.0), 2)
alipay_revenue = (await db.execute(
select(func.coalesce(func.sum(PaymentOrder.amount), 0)).where(
@@ -1904,21 +1984,31 @@ async def get_stats(
)
)).scalar() or 0
last_period_credits_consumed = (await db.execute(
select(func.coalesce(func.sum(func.abs(CreditRecord.amount)), 0)).where(
CreditRecord.type == "consume",
real_credit_charge_filter,
last_period_net_row = (await db.execute(
select(
func.coalesce(func.sum(_stats_real_and_hold), 0),
func.coalesce(func.sum(_stats_refund_and_release), 0),
).where(
CreditRecord.type.in_(["consume", "refund"]),
CreditRecord.created_at >= last_period_start,
CreditRecord.created_at <= last_period_end,
)
)).scalar() or 0
)).one()
last_period_credits_consumed = round(
max(float(last_period_net_row[0] or 0) - float(last_period_net_row[1] or 0), 0.0), 2,
)
# ── 每日各模块积分消耗(始终返回选中日期往前7天,便于图表展示)
# created_at 为 timestamptz,数据库 session 时区已是东八区(CST),
# 读取出来的时间值即为北京时间,直接 CAST 成日期即可,无需再 +8 小时。
from sqlalchemy import Date, cast as sa_cast
_day_expr = sa_cast(CreditRecord.created_at, Date)
# 图表固定展示 [date_end - 6天, date_end] 共7天
_chart_end_dt = date_end
_chart_start_dt = _chart_end_dt - timedelta(days=6)
_chart_start_dt = datetime(
_chart_end_dt.year, _chart_end_dt.month, _chart_end_dt.day, 0, 0, 0, 0, tzinfo=CST,
) - timedelta(days=6)
_chart_end_dt_inclusive = _chart_end_dt.replace(hour=23, minute=59, second=59, microsecond=999999)
_inner = (
select(
_day_expr.label('date'),
@@ -1929,7 +2019,7 @@ async def get_stats(
CreditRecord.type == "consume",
real_credit_charge_filter,
CreditRecord.created_at >= _chart_start_dt,
CreditRecord.created_at <= _chart_end_dt,
CreditRecord.created_at <= _chart_end_dt_inclusive,
)
.group_by(_day_expr, CreditRecord.source_module)
.subquery()
@@ -1973,23 +2063,41 @@ async def get_stats(
]
# ── 各团队积分消耗(有团队 vs 无团队,使用流水中的团队快照)
# 净消耗 = (真实消费 charge + 预扣占用 hold) - (真实退款 refund + 预扣释放 hold_release)
# 注意:
# hold(预扣占用):type=consumecharge_action='hold'amount<0 → 加项
# hold_release(预扣释放):type=refundcharge_action='hold_release'amount>0 → 减项(type=refund 天然包含)
# charge(真实扣费):type=consumecharge/NULL → 加项
# refund(真实退款):type=refundrefund/NULL → 减项
_charge_hold_filter = and_(
CreditRecord.type == "consume",
credit_charge_action_filter, # charge / hold / NULL(历史 charge)
)
_charge_hold_expr = case((_charge_hold_filter, func.abs(CreditRecord.amount)), else_=0)
# type=refund = 真实退款 + 预扣释放退回(账本强制 hold_release.type=refund
_refund_release_expr = case((CreditRecord.type == "refund", func.abs(CreditRecord.amount)), else_=0)
team_credit_rows = (await db.execute(
select(
func.coalesce(CreditRecord.team_name_snapshot, '未分配团队').label('team_name'),
CreditRecord.team_id_snapshot.label('team_id'),
func.coalesce(func.sum(func.abs(CreditRecord.amount)), 0).label('credits'),
func.coalesce(func.sum(_charge_hold_expr), 0).label("total_charge_hold"),
func.coalesce(func.sum(_refund_release_expr), 0).label("total_refund_release"),
)
.where(
CreditRecord.type == "consume",
real_credit_charge_filter,
CreditRecord.type.in_(["consume", "refund"]),
CreditRecord.created_at >= date_start,
CreditRecord.created_at <= date_end,
)
.group_by(CreditRecord.team_id_snapshot, CreditRecord.team_name_snapshot)
.order_by(func.coalesce(func.sum(func.abs(CreditRecord.amount)), 0).desc())
# 按"净消耗 = 真实+预扣 - 退款+释放"倒序排序(排行榜)
.order_by((func.coalesce(func.sum(_charge_hold_expr), 0) - func.coalesce(func.sum(_refund_release_expr), 0)).desc())
)).all()
credits_by_team = [
TeamCreditOut(team_name=row.team_name, team_id=row.team_id, credits=float(row.credits or 0))
TeamCreditOut(
team_name=row.team_name,
team_id=row.team_id,
credits=round(max(float(row.total_charge_hold or 0) - float(row.total_refund_release or 0), 0.0), 2),
)
for row in team_credit_rows
]
@@ -2127,6 +2235,8 @@ async def admin_list_generation_records(
status: str | None = Query(None),
engine_id: str | None = Query(None),
include_media_references: bool | None = Query(None),
start_date: str | None = Query(None, description="创建时间起始,格式 YYYY-MM-DD"),
end_date: str | None = Query(None, description="创建时间结束,格式 YYYY-MM-DD"),
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=500),
admin: User = Depends(get_admin_user),
@@ -2149,6 +2259,10 @@ async def admin_list_generation_records(
query = query.where(GenerationRecord.engine_id == engine_id)
if include_media_references is not None:
query = query.where(GenerationRecord.include_media_references.is_(include_media_references))
if start_date:
query = query.where(GenerationRecord.created_at >= datetime.strptime(start_date, "%Y-%m-%d").replace(tzinfo=CST))
if end_date:
query = query.where(GenerationRecord.created_at < (datetime.strptime(end_date, "%Y-%m-%d") + timedelta(days=1)).replace(tzinfo=CST))
# Count total
count_query = (
@@ -2164,6 +2278,10 @@ async def admin_list_generation_records(
count_query = count_query.where(GenerationRecord.engine_id == engine_id)
if include_media_references is not None:
count_query = count_query.where(GenerationRecord.include_media_references.is_(include_media_references))
if start_date:
count_query = count_query.where(GenerationRecord.created_at >= datetime.strptime(start_date, "%Y-%m-%d").replace(tzinfo=CST))
if end_date:
count_query = count_query.where(GenerationRecord.created_at < (datetime.strptime(end_date, "%Y-%m-%d") + timedelta(days=1)).replace(tzinfo=CST))
total_result = await db.execute(count_query)
total = total_result.scalar() or 0
@@ -2404,6 +2522,118 @@ async def upload_login_video(
return {"url": url}
# ── Payment Stats ────────────────────────────────────────
# ── Invoice Management ───────────────────────────────────
@router.get("/invoices")
async def admin_list_invoices(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=500),
status: str | None = Query(None),
phone: str | None = Query(None, description="按用户手机号模糊搜索"),
start_date: str | None = Query(None),
end_date: str | None = Query(None),
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""后台发票列表(分页+筛选)。"""
from app.services.invoice import get_admin_invoices
items, total = await get_admin_invoices(
db, page, page_size,
status_filter=status,
phone=phone,
start_date=start_date,
end_date=end_date,
)
return {"items": items, "total": total, "page": page, "page_size": page_size}
@router.get("/invoices/{invoice_id}")
async def admin_get_invoice(
invoice_id: str,
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""后台获取发票详情(含关联订单)。"""
from app.services.invoice import get_invoice_with_orders
detail = await get_invoice_with_orders(db, invoice_id)
if not detail:
raise HTTPException(status_code=404, detail="发票不存在")
invoice = detail["invoice"]
orders = detail["orders"]
return {
"id": invoice.id,
"invoiceNo": invoice.invoice_no,
"userId": invoice.user_id,
"headerType": invoice.header_type,
"headerName": invoice.header_name,
"headerTaxNo": invoice.header_tax_no,
"headerRegisterAddress": invoice.header_register_address,
"headerRegisterPhone": invoice.header_register_phone,
"headerBankName": invoice.header_bank_name,
"headerBankAccount": invoice.header_bank_account,
"email": invoice.email,
"totalAmount": round(float(invoice.total_amount), 2),
"totalCredits": round(float(invoice.total_credits), 2),
"status": invoice.status,
"failureReason": invoice.failure_reason,
"issuedAt": invoice.issued_at.isoformat() if invoice.issued_at else None,
"createdAt": invoice.created_at.isoformat() if invoice.created_at else None,
"updatedAt": invoice.updated_at.isoformat() if invoice.updated_at else None,
"orders": [
{
"id": o.id,
"orderNo": o.order_no,
"amount": round(float(o.amount), 2),
"credits": round(float(o.credits), 2),
}
for o in orders
],
}
@router.put("/invoices/{invoice_id}/status")
async def admin_update_invoice_status(
invoice_id: str,
req: InvoiceStatusUpdateRequest,
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""更新发票状态(success/failed)。"""
from app.services.invoice import update_invoice_status
invoice, old_status = await update_invoice_status(db, invoice_id, req, admin.id)
await db.flush()
await log_operation(
db,
admin.id,
admin.username,
f"发票状态变更: {invoice.invoice_no} {old_status}{req.status}",
"PUT",
f"/admin/invoices/{invoice_id}/status",
detail=json.dumps(
{
"invoice_id": invoice_id,
"invoice_no": invoice.invoice_no,
"old_status": old_status,
"new_status": req.status,
"failure_reason": req.failure_reason,
},
ensure_ascii=False,
),
)
await db.commit()
return {
"id": invoice.id,
"invoiceNo": invoice.invoice_no,
"status": invoice.status,
"failureReason": invoice.failure_reason,
"issuedAt": invoice.issued_at.isoformat() if invoice.issued_at else None,
}
+4 -1
View File
@@ -351,7 +351,7 @@ async def get_site_info(db: AsyncSession = Depends(get_db)):
"""Public endpoint returning site name, logo, agreement and copyright info."""
result = await db.execute(
select(SystemConfig).where(SystemConfig.key.in_([
"site_name", "site_logo", "user_agreement_privacy_url", "site_copyright", "operation_manual", "login_bg_video"
"site_name", "site_logo", "user_agreement_privacy_url", "site_copyright", "operation_manual", "login_bg_video", "optimize_hold_credits", "site_banner", "site_banner_version"
]))
)
configs = result.scalars().all()
@@ -375,6 +375,9 @@ async def get_site_info(db: AsyncSession = Depends(get_db)):
"site_copyright": info.get("site_copyright", "© 2026 智创 版权所有"),
"operation_manual": info.get("operation_manual", ""),
"login_bg_video": to_full_url(info.get("login_bg_video")) if info.get("login_bg_video") else "",
"optimize_hold_credits": int(info.get("optimize_hold_credits") or 5),
"site_banner": info.get("site_banner", ""),
"site_banner_version": int(info.get("site_banner_version") or 0),
}
@@ -0,0 +1,96 @@
import logging
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_db, get_current_user
from app.models.user import User
from app.schemas.invoice import InvoiceHeaderCreate, InvoiceHeaderOut, InvoiceHeaderUpdate
from app.services.invoice_header import (
create_header,
delete_header,
get_user_headers,
set_default_header,
update_header,
)
logger = logging.getLogger("videogen")
router = APIRouter(prefix="/invoice-headers", tags=["invoice-headers"])
def _header_to_out(header) -> dict:
return {
"id": header.id,
"user_id": header.user_id,
"type": header.type,
"name": header.name,
"tax_no": header.tax_no,
"register_address": header.register_address,
"register_phone": header.register_phone,
"bank_name": header.bank_name,
"bank_account": header.bank_account,
"email": header.email,
"is_default": header.is_default,
"created_at": header.created_at.isoformat() if header.created_at else None,
"updated_at": header.updated_at.isoformat() if header.updated_at else None,
}
@router.post("", response_model=InvoiceHeaderOut)
async def create_invoice_header(
req: InvoiceHeaderCreate,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""创建发票抬头。"""
header = await create_header(db, current_user.id, req)
await db.commit()
return _header_to_out(header)
@router.get("")
async def list_invoice_headers(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""获取当前用户的所有发票抬头。"""
headers = await get_user_headers(db, current_user.id)
return {"items": [_header_to_out(h) for h in headers]}
@router.put("/{header_id}", response_model=InvoiceHeaderOut)
async def update_invoice_header(
header_id: str,
req: InvoiceHeaderUpdate,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""更新发票抬头。"""
header = await update_header(db, header_id, current_user.id, req)
await db.commit()
return _header_to_out(header)
@router.delete("/{header_id}")
async def delete_invoice_header(
header_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""删除发票抬头。"""
await delete_header(db, header_id, current_user.id)
await db.commit()
return {"success": True}
@router.put("/{header_id}/set-default", response_model=InvoiceHeaderOut)
async def set_default_invoice_header(
header_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""设置默认发票抬头。"""
header = await set_default_header(db, header_id, current_user.id)
await db.commit()
return _header_to_out(header)
+110
View File
@@ -0,0 +1,110 @@
import logging
from fastapi import APIRouter, Depends, HTTPException, Query, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_db, get_current_user
from app.models.invoice import Invoice, InvoiceOrder
from app.models.user import User
from app.schemas.invoice import InvoiceCreateRequest, InvoiceOut, InvoiceOrderOut
from app.services.invoice import (
create_invoice,
get_user_invoices,
get_invoice_by_id,
get_invoice_with_orders,
)
logger = logging.getLogger("videogen")
router = APIRouter(prefix="/invoices", tags=["invoices"])
@router.post("", response_model=InvoiceOut)
async def create_invoice_endpoint(
req: InvoiceCreateRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""创建发票申请。"""
invoice = await create_invoice(db, current_user.id, req)
await db.commit()
# 重新查询以获取关联订单
detail = await get_invoice_with_orders(db, invoice.id)
return _invoice_to_out(detail["invoice"], detail["orders"])
@router.get("")
async def list_invoices(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""获取当前用户的发票列表(分页)。"""
invoices, total = await get_user_invoices(db, current_user.id, page, page_size)
# 加载每个发票的关联订单
items = []
for inv in invoices:
result = await db.execute(
select(InvoiceOrder).where(InvoiceOrder.invoice_id == inv.id)
)
orders = result.scalars().all()
items.append(_invoice_to_out(inv, list(orders)))
return {"items": items, "total": total, "page": page, "page_size": page_size}
@router.get("/{invoice_id}")
async def get_invoice(
invoice_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""获取发票详情(含关联订单)。"""
detail = await get_invoice_with_orders(db, invoice_id)
if not detail:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="发票不存在")
invoice = detail["invoice"]
if invoice.user_id != current_user.id:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="无权查看该发票")
return _invoice_to_out(invoice, detail["orders"])
def _invoice_to_out(invoice: Invoice, orders: list[InvoiceOrder]) -> dict:
"""将 Invoice ORM 对象转换为响应 dict。"""
return {
"id": invoice.id,
"user_id": invoice.user_id,
"invoice_no": invoice.invoice_no,
"header_type": invoice.header_type,
"header_name": invoice.header_name,
"header_tax_no": invoice.header_tax_no,
"header_register_address": invoice.header_register_address,
"header_register_phone": invoice.header_register_phone,
"header_bank_name": invoice.header_bank_name,
"header_bank_account": invoice.header_bank_account,
"email": invoice.email,
"total_amount": round(float(invoice.total_amount), 2),
"total_credits": round(float(invoice.total_credits), 2),
"status": invoice.status,
"failure_reason": invoice.failure_reason,
"issued_at": invoice.issued_at.isoformat() if invoice.issued_at else None,
"created_at": invoice.created_at.isoformat() if invoice.created_at else None,
"updated_at": invoice.updated_at.isoformat() if invoice.updated_at else None,
"orders": [
{
"id": o.id,
"invoice_id": o.invoice_id,
"order_id": o.order_id,
"order_no": o.order_no,
"amount": round(float(o.amount), 2),
"credits": round(float(o.credits), 2),
}
for o in orders
],
}
+52 -3
View File
@@ -289,17 +289,36 @@ async def alipay_callback(request: Request, db: AsyncSession = Depends(get_db)):
async def list_orders(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
status_filter: str | None = Query(None, description="按状态筛选: pending/paid/refunded/failed/cancelled"),
start_date: str | None = Query(None, description="创建时间起始,格式 YYYY-MM-DD"),
end_date: str | None = Query(None, description="创建时间结束,格式 YYYY-MM-DD"),
invoice_mode: bool = Query(False, description="开票模式:仅返回已支付订单"),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
from app.services.payment import _check_and_expire_order
from datetime import datetime, timezone, timedelta
count_query = select(func.count(PaymentOrder.id)).where(PaymentOrder.user_id == current_user.id)
# 构建筛选条件
conditions = [PaymentOrder.user_id == current_user.id]
if status_filter:
conditions.append(PaymentOrder.status == status_filter)
if invoice_mode:
conditions.append(PaymentOrder.status == "paid")
if start_date:
start_dt = datetime.strptime(start_date, "%Y-%m-%d").replace(tzinfo=timezone.utc)
conditions.append(PaymentOrder.created_at >= start_dt)
if end_date:
end_dt = (datetime.strptime(end_date, "%Y-%m-%d") + timedelta(days=1)).replace(tzinfo=timezone.utc)
conditions.append(PaymentOrder.created_at < end_dt)
# 统计总数
count_query = select(func.count(PaymentOrder.id)).where(*conditions)
total = (await db.execute(count_query)).scalar() or 0
result = await db.execute(
select(PaymentOrder)
.where(PaymentOrder.user_id == current_user.id)
.where(*conditions)
.order_by(PaymentOrder.created_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
@@ -308,7 +327,37 @@ async def list_orders(
for o in orders:
await _check_and_expire_order(db, o)
return {"items": [PaymentOrderOut.model_validate(o) for o in orders], "total": total}
# 开票模式:附带订单占用状态
items = []
if invoice_mode:
# 收集当前页订单ID
order_ids = [o.id for o in orders]
# 查询这些订单是否已被占用
from app.models.invoice import Invoice, InvoiceOrder
occupied_map: dict[str, str] = {}
if order_ids:
occ_result = await db.execute(
select(InvoiceOrder.order_id, Invoice.invoice_no)
.join(Invoice, InvoiceOrder.invoice_id == Invoice.id)
.where(
InvoiceOrder.order_id.in_(order_ids),
Invoice.status.in_(["processing", "success"]),
)
)
for row in occ_result.all():
occupied_map[row.order_id] = row.invoice_no
for o in orders:
item = PaymentOrderOut.model_validate(o)
item_dict = item.model_dump()
item_dict["is_occupied"] = o.id in occupied_map
item_dict["occupied_by"] = occupied_map.get(o.id)
items.append(item_dict)
else:
for o in orders:
item = PaymentOrderOut.model_validate(o)
items.append(item.model_dump())
return {"items": items, "total": total}
@router.get("/orders/{order_no}", response_model=PaymentOrderOut)
+29 -3
View File
@@ -364,15 +364,41 @@ async def export_team_credit_records(
end_date=end_date,
)
# 生成 CSV(兼容 Excel 打开)
# 生成 CSV(兼容 Excel 打开UTF-8 BOM
import csv
import io
from datetime import datetime as _dt
def _format_dt(val):
if val is None:
return "-"
return str(datetime.fromtimestamp(val).strftime("%Y-%m-%d %H:%M:%S"))
try:
# 情况 1:已经是 datetime
if isinstance(val, _dt):
dt = val
elif isinstance(val, (int, float)):
# 情况 2:Unix 时间戳(极少,兼容旧代码)
dt = _dt.fromtimestamp(val)
elif isinstance(val, str):
# 情况 3ISO 字符串(admin_credit_record_service._iso 返回的格式)
s = val.strip()
if s.endswith("Z"):
s = s[:-1] + "+00:00"
try:
dt = _dt.fromisoformat(s)
except ValueError:
# 兼容旧格式 YYYY-MM-DD HH:MM:SS
dt = _dt.strptime(s, "%Y-%m-%d %H:%M:%S")
else:
return str(val)
# 统一转东八区展示
if getattr(dt, "tzinfo", None) is None:
dt = dt.replace(tzinfo=CST)
else:
dt = dt.astimezone(CST)
return dt.strftime("%Y-%m-%d %H:%M:%S")
except Exception: # noqa: BLE001
return str(val) if val else "-"
output = io.StringIO()
writer = csv.writer(output)
+12
View File
@@ -0,0 +1,12 @@
from fastapi import APIRouter
from app.api.v3.videos import router as videos_router
from app.api.v3.images import router as images_router
from app.api.v3.models import router as models_router
from app.api.v3.virtual_portrait import router as virtual_portrait_router
api_router_v3 = APIRouter()
api_router_v3.include_router(models_router)
api_router_v3.include_router(videos_router)
api_router_v3.include_router(images_router)
api_router_v3.include_router(virtual_portrait_router)
+14
View File
@@ -0,0 +1,14 @@
from pydantic import BaseModel
class ApiError(BaseModel):
"""API 错误详情。"""
code: str
message: str
class ApiErrorResponse(BaseModel):
"""API 错误响应(旧格式,保留兼容)。"""
error: ApiError
+55
View File
@@ -0,0 +1,55 @@
import logging
import time
from fastapi import APIRouter, Depends, HTTPException, status
from fastapi.responses import JSONResponse
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_db
from app.schemas.api_v3.image import (
ApiImageGenerateRequest,
ApiImageGenerateResponse,
)
from app.services.api_v3 import auth_service, generation_service
logger = logging.getLogger("videogen")
router = APIRouter(prefix="/images", tags=["api-v3-images"])
@router.post(
"",
summary="生成图片",
description="同步生成图片,等待完成后直接返回结果",
)
async def generate_image(
req: ApiImageGenerateRequest,
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
) -> JSONResponse:
"""同步生成图片。"""
start_time = time.perf_counter()
try:
result = await generation_service.generate_image_sync(
db=db,
key=key_context.api_key,
callable_models=key_context.callable_models,
req=req,
start_time=start_time,
)
data = result.model_dump()
# 处理 datetime 序列化
if data.get("created"):
data["created"] = data["created"] if isinstance(data["created"], int) else int(data["created"])
return JSONResponse(
content={"code": 0, "data": data, "message": "ok"},
status_code=200,
)
except HTTPException:
raise
except Exception as exc:
logger.exception("API image generation failed")
raise HTTPException(
status_code=status.HTTP_504_GATEWAY_TIMEOUT,
detail=f"图片生成失败: {str(exc)[:200]}",
)
+119
View File
@@ -0,0 +1,119 @@
import json
import logging
from fastapi import APIRouter, Depends
from fastapi.responses import JSONResponse
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_db
from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine
from app.schemas.api_v3.model import ApiModelInfo, ApiModelsResponse
from app.services.api_v3 import auth_service
from app.services.api_v3.pricing_service import get_priced_models
logger = logging.getLogger("videogen")
router = APIRouter(prefix="/models", tags=["api-v3-models"])
@router.get(
"",
summary="获取可用模型列表",
description="获取当前 API Key 可调用的所有视频和图片模型(仅返回已配置价格的模型)",
)
async def list_models(
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
) -> JSONResponse:
"""获取当前 API Key 可用的模型列表。"""
models: list[ApiModelInfo] = []
# 获取所有已配置价格的引擎 ID 集合
priced_engine_ids = await get_priced_models(db)
# 获取 API Key 的白名单引擎 ID 集合
allowed_engine_ids = {m.get("engine_id", "") for m in key_context.callable_models} if key_context.callable_models else set()
# 确定要返回的引擎 ID 列表
target_engine_ids = priced_engine_ids if not allowed_engine_ids else (allowed_engine_ids & priced_engine_ids)
# 构建引擎信息映射
engine_info_map = {m.get("engine_id", ""): m for m in key_context.callable_models}
for engine_id in target_engine_ids:
engine_type = engine_info_map.get(engine_id, {}).get("engine_type", "")
model_name = engine_info_map.get(engine_id, {}).get("model_name", "")
# 如果没有从白名单获取到类型,尝试从数据库加载
if not engine_type:
video_result = await db.execute(
select(VideoEngine).where(VideoEngine.id == engine_id, VideoEngine.deleted_at.is_(None)).limit(1)
)
if video_result.scalar_one_or_none():
engine_type = "video"
else:
image_result = await db.execute(
select(ImageEngine).where(ImageEngine.id == engine_id, ImageEngine.deleted_at.is_(None)).limit(1)
)
if image_result.scalar_one_or_none():
engine_type = "image"
# 加载引擎详情
supported_ratios = None
supported_resolutions = None
supported_durations = None
supported_sizes = None
try:
if engine_type == "video":
result = await db.execute(
select(VideoEngine).where(VideoEngine.id == engine_id, VideoEngine.deleted_at.is_(None)).limit(1)
)
engine = result.scalar_one_or_none()
if engine:
if not model_name:
model_name = engine.model_name
supported_ratios = _parse_json_list(engine.supported_ratios)
supported_resolutions = _parse_json_list(engine.supported_resolutions)
supported_durations = _parse_json_list(engine.supported_durations)
elif engine_type == "image":
result = await db.execute(
select(ImageEngine).where(ImageEngine.id == engine_id, ImageEngine.deleted_at.is_(None)).limit(1)
)
engine = result.scalar_one_or_none()
if engine:
if not model_name:
model_name = engine.model_name
supported_sizes = _parse_json_list(engine.supported_sizes)
except Exception:
pass
info = ApiModelInfo(
model=model_name,
engine_type=engine_type,
engine_id=engine_id,
supported_ratios=supported_ratios,
supported_resolutions=supported_resolutions,
supported_durations=supported_durations,
supported_sizes=supported_sizes,
)
models.append(info)
return JSONResponse(
content={"code": 0, "data": {"models": [m.model_dump() for m in models]}, "message": "ok"},
status_code=200,
)
def _parse_json_list(value: str | None) -> list[str | int] | None:
"""解析 JSON 列表字段。"""
if not value:
return None
try:
parsed = json.loads(value)
return parsed if isinstance(parsed, list) else None
except (json.JSONDecodeError, TypeError):
return None
+167
View File
@@ -0,0 +1,167 @@
import logging
import time
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_db
from app.models.api.api_generation_task import ApiGenerationTask
from app.schemas.api_v3.video import (
ApiVideoCreateRequest,
ApiVideoCreateResponse,
ApiVideoStatusResponse,
)
from app.services.api_v3 import auth_service, generation_service, task_service
from app.services.resource_signed_url_service import build_resource_signed_url
logger = logging.getLogger("videogen")
router = APIRouter(prefix="/videos", tags=["api-v3-videos"])
async def _validate_request(
db: AsyncSession,
key_context: auth_service.ApiKeyContext,
req: ApiVideoCreateRequest,
) -> ApiGenerationTask | None:
"""请求层校验:参数、权限、幂等性。
Returns:
None = 校验通过,继续创建
ApiGenerationTask = 幂等请求,返回已有任务
"""
# 模型权限校验
allowed_model_names = {m.get("model_name", "") for m in key_context.callable_models}
if allowed_model_names and req.model not in allowed_model_names:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=f"无权使用模型 {req.model}",
)
# 幂等性检查
if req.idempotency_key:
result = await db.execute(
select(ApiGenerationTask).where(
ApiGenerationTask.api_key_id == key_context.api_key.id,
ApiGenerationTask.external_idempotency_key == req.idempotency_key,
ApiGenerationTask.deleted_at.is_(None),
).limit(1)
)
existing_task = result.scalar_one_or_none()
if existing_task:
logger.info(
"Idempotent request: returning existing task %s for key=%s",
existing_task.id, req.idempotency_key,
)
return existing_task
return None
def _map_status(internal_status: str) -> str:
"""将内部状态映射为 API 状态。"""
status_map = {
"pending": "queued",
"queued": "queued",
"generating": "running",
"processing": "running",
"completed": "succeeded",
"failed": "failed",
"timeout": "expired",
}
return status_map.get(internal_status, internal_status)
@router.post(
"",
response_model=ApiVideoCreateResponse,
summary="创建视频生成任务",
)
async def create_video(
req: ApiVideoCreateRequest,
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
) -> ApiVideoCreateResponse:
"""创建视频生成任务(异步)。
幂等性说明:如果 idempotency_key 已存在,直接返回已有任务 ID(不会重复创建)。
"""
try:
# 路由层校验:权限、幂等性
existing_task = await _validate_request(db, key_context, req)
if existing_task:
logger.info(
"Idempotent request: returning existing task %s for key=%s",
existing_task.id, req.idempotency_key,
)
return ApiVideoCreateResponse(id=f"zc-{existing_task.id}")
# 调用服务层创建任务
result = await generation_service.submit_video_generation(
db=db,
key=key_context.api_key,
callable_models=key_context.callable_models,
req=req,
)
return ApiVideoCreateResponse(id=f"zc-{result.id}")
except HTTPException:
raise
except Exception as exc:
logger.exception("API video creation failed")
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"创建视频任务失败: {str(exc)[:200]}",
)
@router.get(
"/{task_id}",
response_model=ApiVideoStatusResponse,
summary="查询视频任务状态",
)
async def get_video_status(
task_id: str,
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
) -> ApiVideoStatusResponse:
"""查询视频任务状态。"""
# 去掉 zc- 前缀
if task_id.startswith("zc-"):
task_id = task_id[3:]
task = await task_service.get_task(db, task_id, key_context.api_key.id)
if not task:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"任务 {task_id} 不存在或不属于当前 API Key",
)
now = int(time.time())
# 构建 content(成功时返回完整视频URL,包含 BASE_URL
content = None
if task.status == "completed" and task.video_url:
from app.schemas.api_v3.video import ApiVideoContent
from app.config import settings
# 拼接完整 URL
video_url = build_resource_signed_url(task.video_url)
if video_url and not video_url.startswith(("http://", "https://")):
base = settings.BASE_URL.rstrip("/")
if video_url.startswith("/"):
video_url = f"{base}{video_url}"
else:
video_url = f"{base}/{video_url}"
content = ApiVideoContent(video_url=video_url)
return ApiVideoStatusResponse(
id=f"zc-{task.id}",
model=task.model_name,
status=_map_status(task.status),
created_at=int(task.created_at.timestamp()) if task.created_at else now,
updated_at=int(task.updated_at.timestamp()) if task.updated_at else now,
content=content,
duration=task.duration,
ratio=task.aspect_ratio,
resolution=task.resolution,
error=task.error_message if task.status in ("failed", "timeout") else None,
)
@@ -0,0 +1,485 @@
from __future__ import annotations
import json
import logging
from datetime import datetime, timedelta, timezone
from fastapi import APIRouter, Depends, Form, HTTPException, Query
from fastapi.responses import JSONResponse
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_db
from app.enums.upload_resource import UploadResourceTypeEnum # noqa: F401 (内部引用保留)
from app.enums.private_portrait import (
PrivatePortraitAssetStatus,
PrivatePortraitAssetType,
PrivatePortraitProjectStatus,
PrivatePortraitRemoteDeleteStatus,
)
from app.schemas.virtual_portrait_v3 import (
VpV3AssetCreate,
VpV3AssetDeleteOut,
VpV3AssetListOut,
VpV3EnumMeta,
VpV3IdOut,
VpV3ProjectCreate,
VpV3ProjectDeleteOut,
VpV3ProjectListOut,
VpV3ProjectOut,
VpV3ProjectUpdate,
VpV3QuotaConfigOut,
VpV3SelectableAssetListOut,
)
from app.services import virtual_portrait_v3 as vp_v3
from app.services.api_v3 import auth_service
logger = logging.getLogger("videogen")
router = APIRouter(prefix="/virtual-portrait", tags=["api-v3-virtual-portrait"])
API_PREFIX_INFO = """
> **虚拟素材库(V3 中转 API**
>
> - 数据与前台用户私域素材库完全隔离(独立 `vp_v3_*` 表),归属按 API Key 管理
> - 所有接口需要在 Header 中携带 `Authorization: Bearer <API Key>`(或通过 `X-API-Key`,详见鉴权说明)
> - 配额:每个 API Key 需要管理员在后台配置虚拟素材额度(项目数/素材数/存储 MB),默认 0=不可使用
> - 生命周期:上传文件 → 创建素材(异步审核,会自动轮询)→ 状态 Active 后可用于 AI 创作
> - 远端删除遵循「先本地软删 → commit 后投递 Celery 异步任务删火山」模式,API 返回 `remote_delete_status=pending` 表示处理中
""" # noqa: E501
# ---------------------------------------------------------------------------
# 基础 & 配置
# ---------------------------------------------------------------------------
@router.get(
"/config",
response_model=VpV3QuotaConfigOut,
summary="获取虚拟素材库配额配置",
description=(
"返回当前 API Key 的虚拟素材配额上限(项目/素材/存储)和已使用量。"
"任一上限大于 0 表示启用虚拟素材库功能。"
+ API_PREFIX_INFO
),
)
async def get_virtual_portrait_config(
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
quota = await vp_v3.quota_service.get_quota(db, api_key_id=key_context.api_key_id, refresh=True)
enabled = any([
(quota.project_limit or 0) > 0,
(quota.asset_limit or 0) > 0,
(quota.storage_mb_limit or 0) > 0,
])
return VpV3QuotaConfigOut(
project_limit=int(quota.project_limit or 0),
asset_limit=int(quota.asset_limit or 0),
storage_mb_limit=int(quota.storage_mb_limit or 0),
project_used=int(quota.project_used or 0),
asset_used=int(quota.asset_used or 0),
storage_mb_used=float(quota.storage_mb_used or 0),
enabled=bool(enabled),
)
@router.get(
"/enums",
response_model=VpV3EnumMeta,
summary="获取虚拟素材库枚举元数据",
description="返回素材类型、素材状态、项目状态、远端删除状态等枚举说明。",
)
async def get_virtual_portrait_enums(
_: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
):
return VpV3EnumMeta(
asset_type={
PrivatePortraitAssetType.IMAGE.value: "图片素材",
PrivatePortraitAssetType.VIDEO.value: "视频素材",
},
asset_status={
PrivatePortraitAssetStatus.CREATING.value: "创建中/审核中",
PrivatePortraitAssetStatus.ACTIVE.value: "已就绪/可用",
PrivatePortraitAssetStatus.FAILED.value: "失败",
PrivatePortraitAssetStatus.DELETING.value: "删除中",
},
project_status={
PrivatePortraitProjectStatus.CREATING_REMOTE_GROUP.value: "远端组创建中",
PrivatePortraitProjectStatus.ACTIVE.value: "就绪",
PrivatePortraitProjectStatus.CREATE_GROUP_FAILED.value: "远端组创建失败",
PrivatePortraitProjectStatus.DELETING.value: "删除中",
},
remote_delete_status={
PrivatePortraitRemoteDeleteStatus.NONE.value: "未删除",
PrivatePortraitRemoteDeleteStatus.PENDING.value: "待异步删除",
PrivatePortraitRemoteDeleteStatus.PROCESSING.value: "远端删除中",
PrivatePortraitRemoteDeleteStatus.DELETED.value: "远端已删除",
PrivatePortraitRemoteDeleteStatus.FAILED.value: "远端删除失败",
},
)
# ---------------------------------------------------------------------------
# 项目 CRUD
# ---------------------------------------------------------------------------
@router.post(
"/projects",
response_model=VpV3IdOut,
summary="创建虚拟素材项目",
description=(
"在当前 API Key 下创建一个虚拟素材项目(同步调用火山创建远端 AssetGroup)。"
"项目名称 1-100 字符;描述最多 500 字符。"
"创建项目会占用 1 个项目配额,超出上限将返回 403。"
),
)
async def create_virtual_portrait_project(
payload: VpV3ProjectCreate,
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
try:
project = await vp_v3.project_service.create_project(
db, api_key_id=key_context.api_key_id, payload=payload
)
await db.commit()
except HTTPException:
await db.rollback()
raise
except Exception as exc: # noqa: BLE001
await db.rollback()
raise HTTPException(status_code=500, detail=f"创建项目失败:{exc}") from exc
return VpV3IdOut(Id=project.remote_group_id)
@router.get(
"/projects",
response_model=VpV3ProjectListOut,
summary="查询虚拟素材项目列表",
description="按 API Key 分页查询虚拟素材项目。支持项目名称模糊搜索、状态筛选。默认按创建时间倒序。",
)
async def list_virtual_portrait_projects(
page: int = Query(1, ge=1, description="页码,从 1 开始"),
page_size: int = Query(20, ge=1, le=100, description="每页数量 1-100"),
keyword: str | None = Query(None, description="项目名称模糊搜索"),
status: str | None = Query(None, description="项目状态筛选(不传查全部)"),
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
items, total = await vp_v3.project_service.list_projects(
db,
api_key_id=key_context.api_key_id,
page=page,
page_size=page_size,
keyword=keyword,
status=status,
)
return VpV3ProjectListOut(
items=[vp_v3.project_service.project_to_out(it) for it in items],
total=total,
page=page,
page_size=page_size,
)
@router.get(
"/projects/{project_id}",
response_model=VpV3ProjectOut,
summary="获取虚拟素材项目详情",
)
async def get_virtual_portrait_project(
project_id: str,
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
project = await vp_v3.project_service.get_project(
db, api_key_id=key_context.api_key_id, project_id=project_id
)
return vp_v3.project_service.project_to_out(project)
@router.put(
"/projects/{project_id}",
response_model=VpV3ProjectOut,
summary="更新虚拟素材项目",
description="更新虚拟素材项目本地展示信息(名称/描述),不会重新创建火山远端 Group。",
)
async def update_virtual_portrait_project(
project_id: str,
payload: VpV3ProjectUpdate,
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
try:
project = await vp_v3.project_service.update_project(
db, api_key_id=key_context.api_key_id, project_id=project_id, payload=payload
)
await db.commit()
except HTTPException:
await db.rollback()
raise
except Exception as exc: # noqa: BLE001
await db.rollback()
raise HTTPException(status_code=500, detail=f"更新项目失败:{exc}") from exc
return vp_v3.project_service.project_to_out(project)
@router.delete(
"/projects/{project_id}",
response_model=VpV3ProjectDeleteOut,
summary="删除虚拟素材项目",
description=(
"软删虚拟素材项目及其下所有素材。本地 commit 后会投递 Celery 异步任务去删除火山远端 AssetGroup/Asset。"
"返回的 remote_delete_status=pending 表示远端删除处理中(可通过项目详情接口轮询)。"
),
)
async def delete_virtual_portrait_project(
project_id: str,
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
project = await vp_v3.project_service.soft_delete_project(
db, api_key_id=key_context.api_key_id, project_id=project_id
)
project_id_snapshot = project.id
try:
await db.commit()
except Exception as exc: # noqa: BLE001
await db.rollback()
raise HTTPException(status_code=500, detail=f"删除项目失败:{exc}") from exc
# commit 后投递 V3 专属的异步删除任务
try:
from app.tasks.vp_v3_asset_tasks import delete_v3_project_remote_task # type: ignore
delete_v3_project_remote_task.delay(project_id_snapshot)
logger.info("vp_v3 project %s 已投递远端删除任务", project_id_snapshot)
except Exception as exc: # noqa: BLE001
logger.warning("vp_v3 项目删除任务投递失败:project_id=%s err=%s", project_id_snapshot, exc)
return VpV3ProjectDeleteOut(
success=True,
remote_delete_status=project.remote_delete_status or PrivatePortraitRemoteDeleteStatus.PENDING.value,
)
# ---------------------------------------------------------------------------
# 素材 CRUD
# ---------------------------------------------------------------------------
@router.post(
"/projects/{project_id}/assets",
response_model=VpV3IdOut,
summary="创建虚拟素材(提交审核)",
description=(
"在指定项目下创建虚拟素材,提交到火山进行异步审核。\n"
"- source_url:必填,必须是 POST /uploads/image 或 /uploads/video 返回的 url(或 /uploads/* 路径)\n"
"- asset_typeImage/VideoVideo 必须提供 video_duration(秒),最多 60 秒\n"
"- 创建成功后 status=Creating;建议调用方自行轮询 /assets/{id}/sync 或详情接口直到 status=Active\n"
"- 同时会占用 1 份素材配额和文件大小对应的存储配额"
),
)
async def create_virtual_portrait_asset(
project_id: str,
payload: VpV3AssetCreate,
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
try:
project = await vp_v3.project_service.get_project(
db, api_key_id=key_context.api_key_id, project_id=project_id
)
asset = await vp_v3.asset_service.create_asset(
db, api_key_id=key_context.api_key_id, project=project, payload=payload
)
await db.commit()
except HTTPException:
await db.rollback()
raise
except Exception as exc: # noqa: BLE001
await db.rollback()
raise HTTPException(status_code=500, detail=f"创建素材失败:{exc}") from exc
asset_id_snapshot = asset.remote_asset_id
# commit 成功后投递 V3 专属轮询任务
try:
from app.tasks.vp_v3_asset_tasks import poll_v3_asset_status # type: ignore
async_result = poll_v3_asset_status.delay(asset_id_snapshot)
logger.info(
"vp_v3 素材轮询任务投递成功:asset_id=%s celery_task_id=%s",
asset_id_snapshot,
getattr(async_result, "id", None),
)
except Exception as exc: # noqa: BLE001
logger.warning("vp_v3 素材轮询任务投递失败:asset_id=%s err=%s", asset_id_snapshot, exc)
return VpV3IdOut(Id=asset.remote_asset_id)
@router.get(
"/projects/{project_id}/assets",
response_model=VpV3AssetListOut,
summary="查询指定项目下的虚拟素材列表",
description="按项目分页查询素材。可按 status/asset_type 筛选,按素材名称 keyword 模糊搜索。",
)
async def list_virtual_portrait_project_assets(
project_id: str,
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
status: str | None = Query(None, description="素材状态筛选(Creating/Active/Failed/Deleting"),
keyword: str | None = Query(None, description="素材名称模糊搜索"),
asset_type: str | None = Query(None, description="素材类型:Image/Video"),
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
# 先校验项目归属
await vp_v3.project_service.get_project(db, api_key_id=key_context.api_key_id, project_id=project_id)
items, total = await vp_v3.asset_service.list_assets(
db,
api_key_id=key_context.api_key_id,
project_id=project_id,
status=status,
keyword=keyword,
asset_type=asset_type,
page=page,
page_size=page_size,
)
return VpV3AssetListOut(
items=[vp_v3.asset_service.asset_to_out(it) for it in items],
total=total,
page=page,
page_size=page_size,
)
@router.get(
"/assets/{asset_id}",
summary="获取虚拟素材审核详情",
description=(
"返回素材的 moderation_json(火山审核 JSON)。\n"
"- 若素材状态为 Creating(审核中)且 next_poll_at 已到期,内部会自动调火山 GetAsset 同步最新状态。\n"
"- 返回内容为解析后的 JSON 对象。"
),
)
async def get_virtual_portrait_asset(
asset_id: str,
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
# 北京时间(UTC+8)统一基准
_BJ_TZ = timezone(timedelta(hours=8))
def _bj_now() -> datetime:
"""返回当前北京时间(UTC+8naive datetime。"""
return datetime.now(_BJ_TZ).replace(tzinfo=None)
asset = await vp_v3.asset_service.get_asset(db, api_key_id=key_context.api_key_id, asset_id=asset_id)
# 统一为 naive 北京时间比较
def _naive(dt: datetime | None) -> datetime | None:
if dt is None:
return None
return dt.replace(tzinfo=None) if dt.tzinfo is not None else dt
need_sync = (
asset.status == PrivatePortraitAssetStatus.CREATING.value
and asset.remote_asset_id
and (_naive(asset.next_poll_at) is None or _naive(asset.next_poll_at) <= _bj_now())
)
if need_sync:
try:
asset = await vp_v3.asset_service.sync_asset_status(
db, api_key_id=key_context.api_key_id, asset_id=asset_id,
)
await db.commit()
await db.refresh(asset)
except HTTPException:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
raise HTTPException(status_code=500, detail=f"同步素材状态失败:{exc}") from exc
# 只返回 moderation_json 解析后的内容
moderation = None
if asset.moderation_json:
try:
moderation = json.loads(asset.moderation_json)
except (json.JSONDecodeError, TypeError):
moderation = asset.moderation_json
return JSONResponse(content=moderation)
@router.delete(
"/assets/{asset_id}",
response_model=VpV3AssetDeleteOut,
summary="删除虚拟素材",
description=(
"软删虚拟素材。本地 commit 后会投递 Celery 异步任务去删除火山远端 Asset。"
"返回 remote_delete_status=pending 表示处理中(可通过素材详情接口轮询)。"
),
)
async def delete_virtual_portrait_asset(
asset_id: str,
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
asset = await vp_v3.asset_service.soft_delete_asset(
db, api_key_id=key_context.api_key_id, asset_id=asset_id
)
asset_id_snapshot = asset.remote_asset_id
try:
await db.commit()
except Exception as exc: # noqa: BLE001
await db.rollback()
raise HTTPException(status_code=500, detail=f"删除素材失败:{exc}") from exc
# commit 后投递 V3 专属的异步删除任务
try:
from app.tasks.vp_v3_asset_tasks import delete_v3_asset_remote_task # type: ignore
delete_v3_asset_remote_task.delay(asset_id_snapshot)
except Exception as exc: # noqa: BLE001
logger.warning("vp_v3 素材远端删除任务投递失败:asset_id=%s err=%s", asset_id_snapshot, exc)
return VpV3AssetDeleteOut(
success=True,
remote_delete_status=asset.remote_delete_status or PrivatePortraitRemoteDeleteStatus.PENDING.value,
)
# ---------------------------------------------------------------------------
# AI 创作选择器用
# ---------------------------------------------------------------------------
@router.get(
"/selectable-assets",
response_model=VpV3SelectableAssetListOut,
summary="查询可用于 AI 创作的虚拟素材",
description=(
"只返回当前 API Key 虚拟素材库中 status=Active 的图片/视频素材。"
"该接口提供给 AI 创作参考素材选择器使用。"
),
)
async def list_virtual_portrait_selectable_assets(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
project_id: str | None = Query(None, description="按项目筛选(可选)"),
keyword: str | None = Query(None, description="素材名称模糊搜索"),
asset_type: str | None = Query(None, description="素材类型:Image/Video"),
key_context: auth_service.ApiKeyContext = Depends(auth_service.get_api_key_dependency),
db: AsyncSession = Depends(get_db),
):
items, total = await vp_v3.asset_service.list_selectable_assets(
db,
api_key_id=key_context.api_key_id,
project_id=project_id,
keyword=keyword,
asset_type=asset_type,
page=page,
page_size=page_size,
)
return VpV3SelectableAssetListOut(
items=[vp_v3.asset_service.asset_to_selectable(it) for it in items],
total=total,
page=page,
page_size=page_size,
)