1
This commit is contained in:
@@ -0,0 +1,32 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api.v1.auth import router as auth_router
|
||||
from app.api.v1.projects import router as projects_router
|
||||
from app.api.v1.generation import router as generation_router
|
||||
from app.api.v1.credits import router as credits_router
|
||||
from app.api.v1.payments import router as payments_router
|
||||
from app.api.v1.notifications import router as notifications_router
|
||||
from app.api.v1.captcha import router as captcha_router
|
||||
from app.api.v1.admin import router as admin_router
|
||||
from app.api.v1.sms import router as sms_router
|
||||
from app.api.v1.industries import router as industries_router
|
||||
from app.api.v1.menu_configs import router as menu_configs_router
|
||||
from app.api.v1.recharge_packages import router as recharge_packages_router
|
||||
from app.api.v1.video_engines import router as video_engines_router
|
||||
from app.api.v1.image_engines import router as image_engines_router
|
||||
|
||||
api_router = APIRouter()
|
||||
api_router.include_router(auth_router)
|
||||
api_router.include_router(projects_router)
|
||||
api_router.include_router(generation_router)
|
||||
api_router.include_router(credits_router)
|
||||
api_router.include_router(payments_router)
|
||||
api_router.include_router(notifications_router)
|
||||
api_router.include_router(captcha_router)
|
||||
api_router.include_router(admin_router)
|
||||
api_router.include_router(sms_router)
|
||||
api_router.include_router(industries_router)
|
||||
api_router.include_router(menu_configs_router)
|
||||
api_router.include_router(recharge_packages_router)
|
||||
api_router.include_router(video_engines_router)
|
||||
api_router.include_router(image_engines_router)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,204 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.dependencies import get_db, get_current_user
|
||||
from app.models.system_config import SystemConfig
|
||||
from app.models.user import User
|
||||
from app.schemas.auth import LoginRequest, ChangePasswordRequest, RegisterRequest
|
||||
from app.schemas.user import UserOut
|
||||
from app.services.auth import (
|
||||
authenticate_user,
|
||||
create_access_token,
|
||||
decode_access_token,
|
||||
hash_password,
|
||||
verify_password,
|
||||
)
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
router = APIRouter(prefix="/auth", tags=["auth"])
|
||||
|
||||
|
||||
@router.post("/login")
|
||||
async def login(req: LoginRequest, db: AsyncSession = Depends(get_db)):
|
||||
# Validate captcha in production (when SMS_MOCK is false)
|
||||
if not settings.SMS_MOCK:
|
||||
if not req.captcha_token:
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="需要验证码",
|
||||
)
|
||||
token_sub = decode_access_token(req.captcha_token)
|
||||
if not token_sub or not token_sub.startswith("captcha:"):
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="验证码无效或已过期",
|
||||
)
|
||||
|
||||
user = await authenticate_user(db, req.username, req.password)
|
||||
if not user:
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="用户名或密码错误",
|
||||
)
|
||||
|
||||
# Only allow frontend users to login via this endpoint
|
||||
if user.user_type != "frontend":
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="该账号不允许在此登录",
|
||||
)
|
||||
|
||||
user.last_login_at = datetime.now()
|
||||
user.credits = round(user.credits, 2)
|
||||
await db.flush()
|
||||
|
||||
token = create_access_token(user.id, req.remember_me)
|
||||
return {"access_token": token, "token_type": "bearer", "user": UserOut.model_validate(user)}
|
||||
|
||||
|
||||
@router.post("/register")
|
||||
async def register(req: RegisterRequest, db: AsyncSession = Depends(get_db)):
|
||||
"""Register a new user with phone + SMS code + password."""
|
||||
# Verify SMS code
|
||||
from app.services.sms import verify_sms_code
|
||||
ok = await verify_sms_code(req.phone, req.code)
|
||||
if not ok:
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="验证码错误或已过期",
|
||||
)
|
||||
|
||||
# Check if phone already registered
|
||||
existing = await db.execute(select(User).where(User.phone == req.phone))
|
||||
if existing.scalar_one_or_none():
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="该手机号已注册",
|
||||
)
|
||||
|
||||
# Create user with default username: 用户 + last 4 digits of phone
|
||||
import random
|
||||
username = f"用户{req.phone[-4:]}"
|
||||
existing_name = await db.execute(select(User).where(User.username == username))
|
||||
if existing_name.scalar_one_or_none():
|
||||
username = f"用户{req.phone[-4:]}{random.randint(10, 99)}"
|
||||
|
||||
user = User(
|
||||
id=generate_id(),
|
||||
username=username,
|
||||
phone=req.phone,
|
||||
hashed_password=hash_password(req.password),
|
||||
credits=100,
|
||||
is_admin=False,
|
||||
user_type="frontend",
|
||||
)
|
||||
db.add(user)
|
||||
await db.flush()
|
||||
|
||||
# Assign default menus to new user
|
||||
from app.models.menu_config import MenuConfig
|
||||
result = await db.execute(
|
||||
select(MenuConfig).where(
|
||||
MenuConfig.is_default == True,
|
||||
MenuConfig.is_active == True,
|
||||
MenuConfig.menu_target.in_(["frontend", "both"]),
|
||||
MenuConfig.menu_type == "page",
|
||||
)
|
||||
)
|
||||
default_menus = result.scalars().all()
|
||||
if default_menus:
|
||||
user.allowed_menus = [m.path for m in default_menus if m.path]
|
||||
|
||||
user.credits = round(user.credits, 2)
|
||||
token = create_access_token(user.id)
|
||||
return {"access_token": token, "token_type": "bearer", "user": UserOut.model_validate(user)}
|
||||
|
||||
|
||||
@router.post("/logout")
|
||||
async def logout(current_user: User = Depends(get_current_user)):
|
||||
return {"message": "ok"}
|
||||
|
||||
|
||||
@router.get("/me", response_model=UserOut)
|
||||
async def get_me(current_user: User = Depends(get_current_user)):
|
||||
current_user.credits = round(current_user.credits, 2)
|
||||
return current_user
|
||||
|
||||
|
||||
@router.post("/change-password")
|
||||
async def change_password(
|
||||
req: ChangePasswordRequest,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
if not verify_password(req.old_password, current_user.hashed_password):
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="原密码错误",
|
||||
)
|
||||
|
||||
if len(req.new_password) < 6:
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="密码至少6位",
|
||||
)
|
||||
|
||||
current_user.hashed_password = hash_password(req.new_password)
|
||||
await db.flush()
|
||||
return {"message": "密码修改成功"}
|
||||
|
||||
|
||||
@router.get("/site-info")
|
||||
async def get_site_info(db: AsyncSession = Depends(get_db)):
|
||||
"""Public endpoint returning site name and logo."""
|
||||
result = await db.execute(
|
||||
select(SystemConfig).where(SystemConfig.key.in_([
|
||||
"site_name", "site_logo", "user_agreement_url", "privacy_policy_url"
|
||||
]))
|
||||
)
|
||||
configs = result.scalars().all()
|
||||
info = {c.key: c.value for c in configs}
|
||||
return {
|
||||
"site_name": info.get("site_name", "VideoGen.AI"),
|
||||
"site_logo": info.get("site_logo", ""),
|
||||
"user_agreement_url": info.get("user_agreement_url", ""),
|
||||
"privacy_policy_url": info.get("privacy_policy_url", ""),
|
||||
}
|
||||
|
||||
|
||||
@router.post("/admin-login")
|
||||
async def admin_login(req: LoginRequest, db: AsyncSession = Depends(get_db)):
|
||||
"""Admin-only login endpoint."""
|
||||
user = await authenticate_user(db, req.username, req.password)
|
||||
if not user:
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="用户名或密码错误",
|
||||
)
|
||||
|
||||
if user.user_type != "admin":
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="该账号不是管理员账号",
|
||||
)
|
||||
|
||||
user.last_login_at = datetime.now()
|
||||
user.credits = round(user.credits, 2)
|
||||
await db.flush()
|
||||
|
||||
token = create_access_token(user.id, req.remember_me)
|
||||
return {"access_token": token, "token_type": "bearer", "user": UserOut.model_validate(user)}
|
||||
@@ -0,0 +1,20 @@
|
||||
from fastapi import APIRouter, HTTPException
|
||||
|
||||
from app.schemas.captcha import CaptchaResponse, CaptchaVerifyRequest, CaptchaTokenResponse
|
||||
from app.services.captcha import generate_slider_captcha, verify_slider_captcha
|
||||
|
||||
router = APIRouter(prefix="/captcha", tags=["captcha"])
|
||||
|
||||
|
||||
@router.get("/slider", response_model=CaptchaResponse)
|
||||
async def get_slider_captcha():
|
||||
data = await generate_slider_captcha()
|
||||
return data
|
||||
|
||||
|
||||
@router.post("/verify", response_model=CaptchaTokenResponse)
|
||||
async def verify_captcha(req: CaptchaVerifyRequest):
|
||||
token = await verify_slider_captcha(req.captcha_id, req.x_offset)
|
||||
if not token:
|
||||
raise HTTPException(status_code=400, detail="验证失败,请重试")
|
||||
return CaptchaTokenResponse(token=token)
|
||||
@@ -0,0 +1,41 @@
|
||||
from fastapi import APIRouter, Depends
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.dependencies import get_db, get_current_user
|
||||
from app.models.user import User
|
||||
from app.models.credit_ratio import CreditRatio
|
||||
from app.schemas.credit import CreditBalanceOut, CreditRecordOut
|
||||
from app.schemas.credit_ratio import CreditRatioOut
|
||||
from app.services.credits import get_records
|
||||
from sqlalchemy import select
|
||||
|
||||
router = APIRouter(prefix="/credits", tags=["credits"])
|
||||
|
||||
|
||||
@router.get("", response_model=CreditBalanceOut)
|
||||
async def get_credits(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
records = await get_records(db, current_user.id)
|
||||
return CreditBalanceOut(
|
||||
credits=round(current_user.credits, 2),
|
||||
records=[CreditRecordOut.model_validate(r) for r in records],
|
||||
)
|
||||
|
||||
|
||||
@router.get("/ratios", response_model=dict)
|
||||
async def get_credit_ratios(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(select(CreditRatio))
|
||||
ratios = result.scalars().all()
|
||||
|
||||
grouped = {}
|
||||
for ratio in ratios:
|
||||
if ratio.gen_type not in grouped:
|
||||
grouped[ratio.gen_type] = []
|
||||
grouped[ratio.gen_type].append(CreditRatioOut.model_validate(ratio))
|
||||
|
||||
return grouped
|
||||
@@ -0,0 +1,602 @@
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, UploadFile, File, status
|
||||
from fastapi.responses import RedirectResponse
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.dependencies import get_db, get_current_user
|
||||
from app.models.user import User
|
||||
from app.models.project import Project
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.schemas.generation import (
|
||||
OptimizeParams,
|
||||
GenerateParams,
|
||||
GenerationRecordOut,
|
||||
OptimizeResult,
|
||||
UpdatePromptRequest,
|
||||
GenerationType,
|
||||
DURATIONS,
|
||||
ASPECT_RATIOS,
|
||||
RESOLUTIONS,
|
||||
IMAGE_SIZES,
|
||||
)
|
||||
from app.services.credits import deduct_credits, calc_text_credits, calc_video_credits, calc_image_credits
|
||||
from app.services.llm import optimize_prompt
|
||||
from app.services.video_url import generate_temp_url, validate_and_get_record_id, get_video_stream_url
|
||||
from app.utils.id_gen import generate_id
|
||||
from app.utils.exceptions import InsufficientCreditsError, RecordNotFoundError, InvalidStatusError
|
||||
|
||||
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:
|
||||
try:
|
||||
refs = json.loads(record.media_references)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
refs = None
|
||||
|
||||
error_message = record.error_message
|
||||
if error_message:
|
||||
from app.services.error_codes import ARK_ERRORS
|
||||
import re
|
||||
match = re.search(r"code='([^']+)'", error_message)
|
||||
if match:
|
||||
code = match.group(1)
|
||||
if code in ARK_ERRORS:
|
||||
error_message = ARK_ERRORS[code]
|
||||
else:
|
||||
parts = error_message.split(":")
|
||||
if len(parts) >= 2 and parts[1].strip() in ARK_ERRORS:
|
||||
error_message = ARK_ERRORS[parts[1].strip()]
|
||||
|
||||
return GenerationRecordOut(
|
||||
id=record.id,
|
||||
project_id=record.project_id,
|
||||
project_name=project_name,
|
||||
original_prompt=record.original_prompt,
|
||||
optimized_prompt=record.optimized_prompt,
|
||||
gen_type=record.gen_type,
|
||||
duration=record.duration,
|
||||
aspect_ratio=record.aspect_ratio,
|
||||
resolution=record.resolution,
|
||||
image_size=record.image_size,
|
||||
image_proportion=record.image_proportion,
|
||||
image_px=record.image_px,
|
||||
status=record.status,
|
||||
video_url=record.video_url,
|
||||
image_url=record.image_url,
|
||||
references=refs,
|
||||
text_credits_cost=round(record.text_credits_cost or 0.00, 2),
|
||||
# text_tokens_used=record.text_tokens_used or 0,
|
||||
credits_cost=round(record.credits_cost or 0.00, 2),
|
||||
# video_tokens_used=record.video_tokens_used or 0,
|
||||
# image_tokens_used=record.image_tokens_used or 0,
|
||||
error_message=error_message,
|
||||
created_at=record.created_at,
|
||||
generated_at=record.generated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("", response_model=list[GenerationRecordOut])
|
||||
async def list_records(
|
||||
project_id: str | None = Query(None, alias="project_id"),
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
query = (
|
||||
select(GenerationRecord, Project.name)
|
||||
.join(Project, GenerationRecord.project_id == Project.id)
|
||||
.where(GenerationRecord.user_id == current_user.id)
|
||||
.order_by(GenerationRecord.created_at.desc())
|
||||
)
|
||||
if project_id:
|
||||
query = query.where(GenerationRecord.project_id == project_id)
|
||||
|
||||
result = await db.execute(query)
|
||||
rows = result.all()
|
||||
return [
|
||||
_record_to_out(record, project_name)
|
||||
for record, project_name in rows
|
||||
]
|
||||
|
||||
|
||||
@router.post("/optimize", response_model=OptimizeResult)
|
||||
async def optimize(
|
||||
req: OptimizeParams,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
# Validate parameters based on generation type
|
||||
if req.gen_type == GenerationType.video:
|
||||
if req.duration not in DURATIONS:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长必须为{DURATIONS}秒之一")
|
||||
if not req.duration:
|
||||
raise HTTPException(status_code=400, detail="视频生成需要指定时长")
|
||||
elif req.gen_type == GenerationType.image:
|
||||
if req.image_size not in IMAGE_SIZES:
|
||||
raise HTTPException(status_code=400, detail=f"图片分辨率必须为{IMAGE_SIZES}之一")
|
||||
if not req.image_size:
|
||||
raise HTTPException(status_code=400, detail="图片生成需要指定画面分辨率")
|
||||
|
||||
# Idempotency check: if key provided, return existing record if found
|
||||
if req.idempotency_key:
|
||||
existing = await db.execute(
|
||||
select(GenerationRecord, Project.name)
|
||||
.join(Project, GenerationRecord.project_id == Project.id)
|
||||
.where(
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
GenerationRecord.idempotency_key == req.idempotency_key,
|
||||
GenerationRecord.gen_type == req.gen_type,
|
||||
GenerationRecord.status == "prompt_optimized",
|
||||
)
|
||||
.order_by(GenerationRecord.created_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
row = existing.first()
|
||||
if row:
|
||||
record, project_name = row
|
||||
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),
|
||||
)
|
||||
|
||||
# Check project exists and belongs to user
|
||||
proj_result = await db.execute(
|
||||
select(Project).where(
|
||||
Project.id == req.project_id,
|
||||
Project.user_id == current_user.id,
|
||||
)
|
||||
)
|
||||
project = proj_result.scalar_one_or_none()
|
||||
if not project:
|
||||
raise HTTPException(status_code=404, detail="项目不存在")
|
||||
|
||||
# Optimize prompt via LLM with type-specific context
|
||||
try:
|
||||
optimized, token_usage = await optimize_prompt(
|
||||
db, req.prompt,
|
||||
user_id=current_user.id,
|
||||
industry_key=project.industry,
|
||||
duration=req.duration if req.gen_type == GenerationType.video else None,
|
||||
image_size=req.image_size if req.gen_type == GenerationType.image else None,
|
||||
image_proportion=req.image_proportion if req.gen_type == GenerationType.image else None,
|
||||
image_px=req.image_px if req.gen_type == GenerationType.image else None,
|
||||
references=req.references,
|
||||
gen_type=req.gen_type,
|
||||
)
|
||||
# Create record BEFORE LLM call so it's visible if user refreshes
|
||||
record = GenerationRecord(
|
||||
id=generate_id(),
|
||||
user_id=current_user.id,
|
||||
project_id=req.project_id,
|
||||
original_prompt=req.prompt,
|
||||
gen_type=req.gen_type,
|
||||
duration=req.duration,
|
||||
image_size=req.image_size,
|
||||
image_proportion=req.image_proportion,
|
||||
image_px=req.image_px,
|
||||
status="optimizing",
|
||||
credits_cost=0,
|
||||
text_credits_cost=0,
|
||||
text_tokens_used=0,
|
||||
media_references=json.dumps(req.references) if req.references else None,
|
||||
idempotency_key=req.idempotency_key,
|
||||
)
|
||||
db.add(record)
|
||||
await db.flush()
|
||||
await db.commit()
|
||||
except Exception as e:
|
||||
from app.services.error_codes import extract_error_message
|
||||
# record.status = "failed"
|
||||
# record.error_message = extract_error_message(e, "提示词")
|
||||
# await db.flush()
|
||||
# await db.commit()
|
||||
raise HTTPException(status_code=502, detail=f"AI模型调用失败: {extract_error_message(e, "提示词")}")
|
||||
|
||||
text_credits = await calc_text_credits(
|
||||
db, token_usage["input_tokens"], token_usage["output_tokens"],
|
||||
)
|
||||
|
||||
await deduct_credits(
|
||||
db, current_user.id, text_credits,
|
||||
f"提示词优化 - {project.name}",
|
||||
)
|
||||
|
||||
record.optimized_prompt = optimized
|
||||
record.status = "prompt_optimized"
|
||||
record.text_credits_cost = round(text_credits, 2)
|
||||
record.text_tokens_used = token_usage["total_tokens"]
|
||||
await db.flush()
|
||||
|
||||
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),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/{record_id}/generate")
|
||||
async def generate(
|
||||
record_id: str,
|
||||
req: GenerateParams,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(GenerationRecord, Project.name)
|
||||
.join(Project, GenerationRecord.project_id == Project.id)
|
||||
.where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
)
|
||||
)
|
||||
row = result.first()
|
||||
if not row:
|
||||
raise RecordNotFoundError()
|
||||
|
||||
record, project_name = row
|
||||
if record.status not in ("prompt_optimized", "failed"):
|
||||
raise InvalidStatusError("当前状态不允许生成")
|
||||
|
||||
if record.gen_type == GenerationType.video:
|
||||
# Video generation
|
||||
if req.aspect_ratio not in ASPECT_RATIOS:
|
||||
raise HTTPException(status_code=400, detail="不支持的画面比例")
|
||||
if req.resolution not in RESOLUTIONS:
|
||||
raise HTTPException(status_code=400, detail="不支持的分辨率")
|
||||
|
||||
duration = record.duration or 5
|
||||
video_credits = await calc_video_credits(db, duration, req.resolution)
|
||||
await deduct_credits(
|
||||
db, current_user.id, video_credits,
|
||||
f"视频生成 - {project_name}",
|
||||
related_id=record_id,
|
||||
)
|
||||
|
||||
record.aspect_ratio = req.aspect_ratio
|
||||
record.resolution = req.resolution
|
||||
record.credits_cost = round(video_credits, 2)
|
||||
record.status = "generating"
|
||||
record.error_message = None
|
||||
await db.flush()
|
||||
|
||||
try:
|
||||
from app.services.video_gen import get_active_engine, submit_video_task
|
||||
from app.services.error_codes import extract_error_message
|
||||
from app.services.video_queue import task_queue
|
||||
|
||||
engine = await get_active_engine(db)
|
||||
task_id = await submit_video_task(db, engine, record)
|
||||
record.seedance_task_id = task_id
|
||||
await db.flush()
|
||||
await task_queue.enqueue(record_id)
|
||||
except Exception as e:
|
||||
record.status = "failed"
|
||||
record.error_message = extract_error_message(e, "视频")
|
||||
await db.flush()
|
||||
|
||||
elif record.gen_type == GenerationType.image:
|
||||
# Image generation
|
||||
|
||||
image_credits = await calc_image_credits(db, req.image_size or record.image_size or "2K")
|
||||
await deduct_credits(
|
||||
db, current_user.id, image_credits,
|
||||
f"图片生成 - {project_name}",
|
||||
related_id=record_id,
|
||||
)
|
||||
|
||||
record.image_size = req.image_size or record.image_size or "2K"
|
||||
record.credits_cost = round(image_credits, 2)
|
||||
record.status = "generating"
|
||||
record.error_message = None
|
||||
await db.commit()
|
||||
|
||||
from app.services.video_queue import task_queue
|
||||
await task_queue.enqueue(record_id)
|
||||
|
||||
return _record_to_out(record, project_name)
|
||||
|
||||
|
||||
|
||||
|
||||
@router.post("/{record_id}/retry")
|
||||
async def retry_generation(
|
||||
record_id: str,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(GenerationRecord, Project.name)
|
||||
.join(Project, GenerationRecord.project_id == Project.id)
|
||||
.where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
)
|
||||
)
|
||||
row = result.first()
|
||||
if not row:
|
||||
raise RecordNotFoundError()
|
||||
|
||||
record, project_name = row
|
||||
if record.status != "failed":
|
||||
raise InvalidStatusError("只有失败的记录可以重试")
|
||||
|
||||
# Re-deduct video credits for retry
|
||||
if record.duration and record.resolution:
|
||||
video_credits = await calc_video_credits(db, record.duration, record.resolution)
|
||||
await deduct_credits(
|
||||
db, current_user.id, video_credits,
|
||||
f"视频重试 - {project_name}",
|
||||
related_id=record_id,
|
||||
)
|
||||
record.credits_cost = round((record.credits_cost or 0) + video_credits, 2)
|
||||
|
||||
record.status = "generating"
|
||||
record.error_message = None
|
||||
await db.flush()
|
||||
|
||||
try:
|
||||
from app.services.video_gen import get_active_engine, submit_video_task, extract_error_message
|
||||
from app.services.video_queue import task_queue
|
||||
|
||||
engine = await get_active_engine(db)
|
||||
task_id = await submit_video_task(db, engine, record)
|
||||
record.seedance_task_id = task_id
|
||||
await db.flush()
|
||||
await task_queue.enqueue(record_id)
|
||||
except Exception as e:
|
||||
record.status = "failed"
|
||||
record.error_message = extract_error_message(e)
|
||||
await db.flush()
|
||||
|
||||
return _record_to_out(record, project_name)
|
||||
|
||||
|
||||
@router.put("/{record_id}/prompt")
|
||||
async def update_prompt(
|
||||
record_id: str,
|
||||
req: UpdatePromptRequest,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(GenerationRecord).where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
)
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
if not record:
|
||||
raise RecordNotFoundError()
|
||||
if record.status != "prompt_optimized":
|
||||
raise InvalidStatusError("只有待生成状态可以修改提示词")
|
||||
|
||||
record.optimized_prompt = req.optimized_prompt
|
||||
await db.flush()
|
||||
return {"message": "ok"}
|
||||
|
||||
|
||||
@router.get("/{record_id}/video")
|
||||
async def get_video(
|
||||
record_id: str,
|
||||
token: str = Query(...),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Validate temp token and redirect to video URL."""
|
||||
validated_id = await validate_and_get_record_id(token)
|
||||
if validated_id != record_id:
|
||||
raise HTTPException(status_code=403, detail="无效的视频链接")
|
||||
|
||||
video_url = await get_video_stream_url(db, record_id)
|
||||
if not video_url:
|
||||
raise HTTPException(status_code=404, detail="视频不存在")
|
||||
|
||||
return RedirectResponse(url=video_url)
|
||||
|
||||
|
||||
@router.get("/{record_id}/queue-status")
|
||||
async def get_queue_status(
|
||||
record_id: str,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Get queue position, estimated wait time, and current status for a generation record."""
|
||||
from sqlalchemy import func
|
||||
|
||||
result = await db.execute(
|
||||
select(GenerationRecord).where(
|
||||
GenerationRecord.id == record_id,
|
||||
GenerationRecord.user_id == current_user.id,
|
||||
)
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
if not record:
|
||||
raise RecordNotFoundError()
|
||||
|
||||
queue_position = None
|
||||
estimated_wait_seconds = None
|
||||
|
||||
if record.status == "generating":
|
||||
ahead_result = await db.execute(
|
||||
select(func.count(GenerationRecord.id)).where(
|
||||
GenerationRecord.status == "generating",
|
||||
GenerationRecord.created_at < record.created_at,
|
||||
)
|
||||
)
|
||||
ahead = ahead_result.scalar() or 0
|
||||
queue_position = ahead + 1
|
||||
estimated_wait_seconds = ahead * 60
|
||||
|
||||
return {
|
||||
"record_id": record.id,
|
||||
"status": record.status,
|
||||
"queue_position": queue_position,
|
||||
"estimated_wait_seconds": estimated_wait_seconds,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/callbacks/seedance")
|
||||
async def seedance_callback(request: Request, db: AsyncSession = Depends(get_db)):
|
||||
"""Receive async callback from Seedance API."""
|
||||
data = await request.json()
|
||||
task_id = data.get("id")
|
||||
task_status = data.get("status")
|
||||
|
||||
if not task_id:
|
||||
return {"message": "ignored"}
|
||||
|
||||
result = await db.execute(
|
||||
select(GenerationRecord).where(GenerationRecord.seedance_task_id == task_id)
|
||||
)
|
||||
record = result.scalar_one_or_none()
|
||||
if not record:
|
||||
return {"message": "record not found"}
|
||||
|
||||
if task_status == "succeeded":
|
||||
remote_url = data.get("content", {}).get("video_url", "")
|
||||
record.status = "completed"
|
||||
# Download video to local storage
|
||||
if settings.STORAGE_TYPE == "local" and remote_url:
|
||||
try:
|
||||
from app.services.video_gen import download_video
|
||||
dest = os.path.join(settings.STORAGE_LOCAL_PATH, f"{record.id}.mp4")
|
||||
await download_video(remote_url, dest)
|
||||
record.video_url = f"/videos/{record.id}.mp4"
|
||||
except Exception as e:
|
||||
logger.warning(f"Callback download failed, using remote URL: {e}")
|
||||
record.video_url = remote_url
|
||||
else:
|
||||
record.video_url = remote_url
|
||||
record.generated_at = datetime.now()
|
||||
# Extract video token usage from callback
|
||||
usage = data.get("usage", {})
|
||||
if usage:
|
||||
record.video_tokens_used = usage.get("total_tokens", 0)
|
||||
# Log callback response
|
||||
from app.services.video_gen import _log_video_response
|
||||
_log_video_response(record.id, data)
|
||||
# Notify user
|
||||
from app.services.notification import create_notification
|
||||
from app.api.v1.notifications import push_notification_to_user
|
||||
notif = await create_notification(
|
||||
db, record.user_id, "视频生成完成",
|
||||
"您的视频已生成完成,可以查看了。", "video", record.id,
|
||||
)
|
||||
await push_notification_to_user(record.user_id, notif)
|
||||
elif task_status == "failed":
|
||||
record.status = "failed"
|
||||
record.error_message = data.get("error", "视频生成失败")
|
||||
# Log callback response
|
||||
from app.services.video_gen import _log_video_response
|
||||
_log_video_response(record.id, data, error=record.error_message)
|
||||
# Notify user
|
||||
from app.services.notification import create_notification
|
||||
from app.api.v1.notifications import push_notification_to_user
|
||||
notif = await create_notification(
|
||||
db, record.user_id, "视频生成失败",
|
||||
f"视频生成失败:{record.error_message}", "video", record.id,
|
||||
)
|
||||
await push_notification_to_user(record.user_id, notif)
|
||||
|
||||
await db.flush()
|
||||
return {"message": "ok"}
|
||||
|
||||
|
||||
@router.post("/upload-image")
|
||||
async def upload_image(
|
||||
file: UploadFile = File(...),
|
||||
current_user: User = Depends(get_current_user),
|
||||
gen_type: str = Query("video", description="生成类型:video-视频,image-图片"),
|
||||
):
|
||||
"""Upload an image for generation reference."""
|
||||
import os
|
||||
import uuid
|
||||
from app.config import settings
|
||||
from datetime import datetime
|
||||
|
||||
if not file.content_type or not file.content_type.startswith("image/"):
|
||||
raise HTTPException(status_code=400, detail="仅支持图片文件")
|
||||
|
||||
ext = os.path.splitext(file.filename or ".png")[1] or ".png"
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
safe_name = f"{gen_type}_img_{current_user.id}_{timestamp}_{uuid.uuid4().hex[:8]}{ext}"
|
||||
date_dir = datetime.now().strftime("%Y/%m/%d")
|
||||
dir_path = os.path.join(settings.UPLOAD_LOCAL_PATH, "images", date_dir)
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
file_path = os.path.join(dir_path, safe_name)
|
||||
|
||||
content = await file.read()
|
||||
if len(content) > 10 * 1024 * 1024:
|
||||
raise HTTPException(status_code=400, detail="图片大小不能超过10MB")
|
||||
|
||||
with open(file_path, "wb") as f:
|
||||
f.write(content)
|
||||
|
||||
url = f"/uploads/images/{date_dir}/{safe_name}"
|
||||
return {"url": url, "filename": file.filename or safe_name, "type": "image", "gen_type": gen_type}
|
||||
|
||||
|
||||
@router.post("/upload-video")
|
||||
async def upload_video(
|
||||
file: UploadFile = File(...),
|
||||
current_user: User = Depends(get_current_user),
|
||||
):
|
||||
"""Upload a video for generation reference."""
|
||||
import os
|
||||
import uuid
|
||||
from app.config import settings
|
||||
from datetime import datetime
|
||||
|
||||
if not file.content_type or not file.content_type.startswith("video/"):
|
||||
raise HTTPException(status_code=400, detail="仅支持视频文件")
|
||||
|
||||
ext = os.path.splitext(file.filename or ".mp4")[1] or ".mp4"
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
safe_name = f"video_ref_{current_user.id}_{timestamp}_{uuid.uuid4().hex[:8]}{ext}"
|
||||
date_dir = datetime.now().strftime("%Y/%m/%d")
|
||||
dir_path = os.path.join(settings.UPLOAD_LOCAL_PATH, "videos", date_dir)
|
||||
os.makedirs(dir_path, exist_ok=True)
|
||||
file_path = os.path.join(dir_path, safe_name)
|
||||
|
||||
content = await file.read()
|
||||
if len(content) > 100 * 1024 * 1024:
|
||||
raise HTTPException(status_code=400, detail="视频大小不能超过100MB")
|
||||
|
||||
with open(file_path, "wb") as f:
|
||||
f.write(content)
|
||||
|
||||
url = f"/uploads/videos/{date_dir}/{safe_name}"
|
||||
return {"url": url, "filename": file.filename or safe_name, "type": "video"}
|
||||
|
||||
|
||||
@router.post("/delete-file")
|
||||
async def delete_upload(
|
||||
url: str = Query(..., description="文件URL,如 /uploads/images/2024/01/01/video_img_xxx.png"),
|
||||
current_user: User = Depends(get_current_user),
|
||||
):
|
||||
"""Delete an uploaded file by URL."""
|
||||
import os
|
||||
from app.config import settings
|
||||
|
||||
if not url.startswith("/uploads/"):
|
||||
raise HTTPException(status_code=400, detail="无效的文件路径")
|
||||
|
||||
if current_user.id not in url:
|
||||
raise HTTPException(status_code=403, detail="无权删除此文件")
|
||||
|
||||
file_path = os.path.join(settings.UPLOAD_LOCAL_PATH, url.replace("/uploads/", ""))
|
||||
if os.path.exists(file_path):
|
||||
os.remove(file_path)
|
||||
|
||||
return {"message": "ok"}
|
||||
@@ -0,0 +1,48 @@
|
||||
import json
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from app.dependencies import get_db, get_current_user
|
||||
from app.models.user import User
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.schemas.image_engine import ImageEngineListResponse
|
||||
|
||||
router = APIRouter(prefix="/image-engines", tags=["image-engines"])
|
||||
|
||||
|
||||
@router.get("", response_model=ImageEngineListResponse)
|
||||
async def list_active_engines(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Public endpoint returning active image engine capabilities."""
|
||||
result = await db.execute(
|
||||
select(ImageEngine)
|
||||
.where(ImageEngine.is_active == True)
|
||||
.order_by(ImageEngine.priority.desc())
|
||||
)
|
||||
engines = result.scalars().all()
|
||||
items = []
|
||||
for e in engines:
|
||||
models = []
|
||||
sizes = {}
|
||||
try:
|
||||
models = json.loads(e.supported_models) if e.supported_models else []
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
sizes = json.loads(e.supported_sizes) if e.supported_sizes else {}
|
||||
except Exception:
|
||||
pass
|
||||
items.append({
|
||||
"id": e.id,
|
||||
"name": e.name,
|
||||
"provider": e.provider,
|
||||
"supported_models": models,
|
||||
"supported_sizes": sizes,
|
||||
"default_size": e.default_size,
|
||||
})
|
||||
return {"items": items}
|
||||
@@ -0,0 +1,46 @@
|
||||
import json
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.dependencies import get_db, get_current_user
|
||||
from app.models.user import User
|
||||
from app.models.industry_config import IndustryConfig
|
||||
|
||||
router = APIRouter(tags=["industries"])
|
||||
|
||||
|
||||
@router.get("/industries")
|
||||
async def list_active_industries(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Public endpoint: list all active industries."""
|
||||
result = await db.execute(
|
||||
select(IndustryConfig)
|
||||
.where(IndustryConfig.is_active == True)
|
||||
.order_by(IndustryConfig.sort_order)
|
||||
)
|
||||
industries = result.scalars().all()
|
||||
items = []
|
||||
for ind in industries:
|
||||
skills = []
|
||||
if ind.skills:
|
||||
try:
|
||||
raw = json.loads(ind.skills)
|
||||
if isinstance(raw, list):
|
||||
skills = raw if raw and isinstance(raw[0], dict) else [{"key": s, "label": s} for s in raw]
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
skills = []
|
||||
items.append({
|
||||
"id": ind.id,
|
||||
"key": ind.key,
|
||||
"label": ind.label,
|
||||
"icon": ind.icon or "",
|
||||
"description": ind.description,
|
||||
"skills": skills,
|
||||
"is_active": ind.is_active,
|
||||
"sort_order": ind.sort_order,
|
||||
})
|
||||
return items
|
||||
@@ -0,0 +1,92 @@
|
||||
from fastapi import APIRouter, Depends
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.dependencies import get_db, get_admin_user, get_current_user
|
||||
from app.models.user import User
|
||||
from app.models.menu_config import MenuConfig
|
||||
from app.schemas.menu import MenuConfigCreate, MenuConfigOut
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
router = APIRouter(tags=["menu-configs"])
|
||||
|
||||
|
||||
@router.get("/menu-configs")
|
||||
async def public_list_menu_configs(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Public endpoint: list active menu configs for frontend (frontend + both)."""
|
||||
result = await db.execute(
|
||||
select(MenuConfig)
|
||||
.where(
|
||||
MenuConfig.is_active == True,
|
||||
MenuConfig.menu_target.in_(["frontend", "both"]),
|
||||
)
|
||||
.order_by(MenuConfig.sort_order)
|
||||
)
|
||||
menus = result.scalars().all()
|
||||
return [
|
||||
{
|
||||
"id": m.id, "label": m.label, "path": m.path, "icon": m.icon,
|
||||
"sort_order": m.sort_order, "is_active": m.is_active,
|
||||
"parent_id": m.parent_id, "menu_type": m.menu_type, "menu_target": m.menu_target,
|
||||
"is_default": m.is_default,
|
||||
}
|
||||
for m in menus
|
||||
]
|
||||
|
||||
|
||||
@router.get("/admin/menu-configs", response_model=list[MenuConfigOut])
|
||||
async def admin_list_menu_configs(
|
||||
admin: User = Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(select(MenuConfig).order_by(MenuConfig.sort_order))
|
||||
return result.scalars().all()
|
||||
|
||||
|
||||
@router.post("/admin/menu-configs", response_model=MenuConfigOut)
|
||||
async def create_menu_config(
|
||||
req: MenuConfigCreate,
|
||||
admin: User = Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
menu = MenuConfig(id=generate_id(), **req.model_dump())
|
||||
db.add(menu)
|
||||
await db.flush()
|
||||
return menu
|
||||
|
||||
|
||||
@router.put("/admin/menu-configs/{menu_id}", response_model=MenuConfigOut)
|
||||
async def update_menu_config(
|
||||
menu_id: str,
|
||||
req: MenuConfigCreate,
|
||||
admin: User = Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(select(MenuConfig).where(MenuConfig.id == menu_id))
|
||||
menu = result.scalar_one_or_none()
|
||||
if not menu:
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="菜单不存在")
|
||||
for k, v in req.model_dump().items():
|
||||
setattr(menu, k, v)
|
||||
await db.flush()
|
||||
return menu
|
||||
|
||||
|
||||
@router.delete("/admin/menu-configs/{menu_id}")
|
||||
async def delete_menu_config(
|
||||
menu_id: str,
|
||||
admin: User = Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(select(MenuConfig).where(MenuConfig.id == menu_id))
|
||||
menu = result.scalar_one_or_none()
|
||||
if not menu:
|
||||
from fastapi import HTTPException, status
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="菜单不存在")
|
||||
await db.delete(menu)
|
||||
await db.flush()
|
||||
return {"message": "ok"}
|
||||
@@ -0,0 +1,142 @@
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime, timezone, timedelta
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, WebSocket, WebSocketDisconnect
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.dependencies import get_db, get_current_user
|
||||
from app.models.user import User
|
||||
from app.models.notification import Notification
|
||||
from app.schemas.notification import NotificationOut, UnreadCountOut
|
||||
from app.services.auth import decode_access_token
|
||||
from app.services.notification import (
|
||||
get_notifications,
|
||||
mark_read,
|
||||
mark_all_read,
|
||||
get_unread_count,
|
||||
create_notification,
|
||||
)
|
||||
|
||||
logger = logging.getLogger("videogen")
|
||||
CST = timezone(timedelta(hours=8))
|
||||
|
||||
|
||||
def _to_local_str(dt):
|
||||
if dt is None:
|
||||
return None
|
||||
if dt.tzinfo and dt.tzinfo.utcoffset(None) == timedelta(0):
|
||||
dt = dt.astimezone(CST)
|
||||
return dt.replace(tzinfo=None).isoformat()
|
||||
|
||||
router = APIRouter(prefix="/notifications", tags=["notifications"])
|
||||
|
||||
# Active WebSocket connections: user_id -> set of WebSockets
|
||||
_active_connections: dict[str, set[WebSocket]] = {}
|
||||
|
||||
|
||||
async def push_notification_to_user(user_id: str, notification: Notification) -> None:
|
||||
"""Push a notification to all active WebSocket connections for a user."""
|
||||
connections = _active_connections.get(user_id, set())
|
||||
if not connections:
|
||||
return
|
||||
message = json.dumps({
|
||||
"id": notification.id,
|
||||
"title": notification.title,
|
||||
"content": notification.content,
|
||||
"type": notification.type,
|
||||
"is_read": notification.is_read,
|
||||
"created_at": _to_local_str(notification.created_at),
|
||||
}, ensure_ascii=False)
|
||||
dead: list[WebSocket] = []
|
||||
for ws in connections:
|
||||
try:
|
||||
await ws.send_text(message)
|
||||
except Exception:
|
||||
dead.append(ws)
|
||||
for ws in dead:
|
||||
connections.discard(ws)
|
||||
|
||||
|
||||
@router.websocket("/ws")
|
||||
async def notifications_ws(websocket: WebSocket):
|
||||
"""Authenticated WebSocket for real-time notifications."""
|
||||
await websocket.accept()
|
||||
|
||||
# Authenticate via token query param or first message
|
||||
token = websocket.query_params.get("token")
|
||||
if not token:
|
||||
try:
|
||||
data = await websocket.receive_text()
|
||||
msg = json.loads(data)
|
||||
token = msg.get("token")
|
||||
except Exception:
|
||||
await websocket.close(code=4001, reason="Authentication required")
|
||||
return
|
||||
|
||||
user_id = decode_access_token(token) if token else None
|
||||
if not user_id or user_id.startswith("captcha:"):
|
||||
await websocket.close(code=4001, reason="Invalid token")
|
||||
return
|
||||
|
||||
# Register connection
|
||||
if user_id not in _active_connections:
|
||||
_active_connections[user_id] = set()
|
||||
_active_connections[user_id].add(websocket)
|
||||
|
||||
try:
|
||||
# Keep connection alive, listen for pings
|
||||
while True:
|
||||
data = await websocket.receive_text()
|
||||
# Echo back for heartbeat
|
||||
if data == "ping":
|
||||
await websocket.send_text("pong")
|
||||
except WebSocketDisconnect:
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
conns = _active_connections.get(user_id)
|
||||
if conns:
|
||||
conns.discard(websocket)
|
||||
if not conns:
|
||||
_active_connections.pop(user_id, None)
|
||||
|
||||
|
||||
@router.get("", response_model=list[NotificationOut])
|
||||
async def list_notifications(
|
||||
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),
|
||||
):
|
||||
items, _ = await get_notifications(db, current_user.id, page, page_size)
|
||||
return items
|
||||
|
||||
|
||||
@router.put("/{notification_id}/read")
|
||||
async def read_notification(
|
||||
notification_id: str,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
await mark_read(db, notification_id, current_user.id)
|
||||
return {"message": "ok"}
|
||||
|
||||
|
||||
@router.put("/read-all")
|
||||
async def read_all(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
await mark_all_read(db, current_user.id)
|
||||
return {"message": "ok"}
|
||||
|
||||
|
||||
@router.get("/unread-count", response_model=UnreadCountOut)
|
||||
async def unread_count(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
count = await get_unread_count(db, current_user.id)
|
||||
return UnreadCountOut(count=count)
|
||||
@@ -0,0 +1,81 @@
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.dependencies import get_db, get_current_user
|
||||
from app.models.user import User
|
||||
from app.models.payment_order import PaymentOrder
|
||||
from app.models.recharge_package import RechargePackage
|
||||
from app.schemas.payment import RechargeRequest, PaymentOrderOut
|
||||
from app.services.payment import create_recharge_order, verify_wechat_callback, verify_alipay_callback, process_payment_success
|
||||
|
||||
router = APIRouter(prefix="/payments", tags=["payments"])
|
||||
|
||||
|
||||
@router.post("/recharge", response_model=PaymentOrderOut)
|
||||
async def recharge(
|
||||
req: RechargeRequest,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(RechargePackage).where(
|
||||
RechargePackage.id == req.plan,
|
||||
RechargePackage.is_active == True,
|
||||
)
|
||||
)
|
||||
pkg = result.scalar_one_or_none()
|
||||
if not pkg:
|
||||
raise HTTPException(status_code=400, detail="无效的套餐")
|
||||
order = await create_recharge_order(
|
||||
db,
|
||||
current_user.id,
|
||||
credits=pkg.credits,
|
||||
price=pkg.price,
|
||||
label=pkg.name,
|
||||
bonus_credits=pkg.bonus_credits,
|
||||
)
|
||||
return order
|
||||
|
||||
|
||||
@router.post("/wechat/callback")
|
||||
async def wechat_callback(request: Request, db: AsyncSession = Depends(get_db)):
|
||||
data = await request.json()
|
||||
if not await verify_wechat_callback(data):
|
||||
raise HTTPException(status_code=400, detail="签名验证失败")
|
||||
order_no = data.get("out_trade_no")
|
||||
result = await db.execute(
|
||||
select(PaymentOrder).where(PaymentOrder.order_no == order_no)
|
||||
)
|
||||
order = result.scalar_one_or_none()
|
||||
if order:
|
||||
await process_payment_success(db, order.id)
|
||||
return {"code": "SUCCESS", "message": "OK"}
|
||||
|
||||
|
||||
@router.post("/alipay/callback")
|
||||
async def alipay_callback(request: Request, db: AsyncSession = Depends(get_db)):
|
||||
data = await request.form()
|
||||
if not await verify_alipay_callback(dict(data)):
|
||||
raise HTTPException(status_code=400, detail="签名验证失败")
|
||||
order_no = data.get("out_trade_no")
|
||||
result = await db.execute(
|
||||
select(PaymentOrder).where(PaymentOrder.order_no == order_no)
|
||||
)
|
||||
order = result.scalar_one_or_none()
|
||||
if order:
|
||||
await process_payment_success(db, order.id)
|
||||
return "success"
|
||||
|
||||
|
||||
@router.get("/orders", response_model=list[PaymentOrderOut])
|
||||
async def list_orders(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(PaymentOrder)
|
||||
.where(PaymentOrder.user_id == current_user.id)
|
||||
.order_by(PaymentOrder.created_at.desc())
|
||||
)
|
||||
return result.scalars().all()
|
||||
@@ -0,0 +1,68 @@
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.dependencies import get_db, get_current_user
|
||||
from app.models.user import User
|
||||
from app.models.project import Project
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.schemas.project import ProjectCreate, ProjectOut
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
router = APIRouter(prefix="/projects", tags=["projects"])
|
||||
|
||||
|
||||
@router.get("", response_model=list[ProjectOut])
|
||||
async def list_projects(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(Project)
|
||||
.where(Project.user_id == current_user.id)
|
||||
.order_by(Project.created_at.desc())
|
||||
)
|
||||
return result.scalars().all()
|
||||
|
||||
|
||||
@router.post("", response_model=ProjectOut)
|
||||
async def create_project(
|
||||
req: ProjectCreate,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
project = Project(
|
||||
id=generate_id(),
|
||||
user_id=current_user.id,
|
||||
name=req.name,
|
||||
industry=req.industry,
|
||||
)
|
||||
db.add(project)
|
||||
await db.flush()
|
||||
return project
|
||||
|
||||
|
||||
@router.delete("/{project_id}")
|
||||
async def delete_project(
|
||||
project_id: str,
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(Project).where(
|
||||
Project.id == project_id,
|
||||
Project.user_id == current_user.id,
|
||||
)
|
||||
)
|
||||
project = result.scalar_one_or_none()
|
||||
if not project:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="项目不存在")
|
||||
|
||||
# Cascade delete generation records
|
||||
from sqlalchemy import delete
|
||||
await db.execute(
|
||||
delete(GenerationRecord).where(GenerationRecord.project_id == project_id)
|
||||
)
|
||||
await db.delete(project)
|
||||
await db.flush()
|
||||
return {"message": "ok"}
|
||||
@@ -0,0 +1,105 @@
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.dependencies import get_db, get_admin_user, get_current_user
|
||||
from app.models.user import User
|
||||
from app.models.recharge_package import RechargePackage
|
||||
from app.schemas.recharge_package import (
|
||||
RechargePackageCreate,
|
||||
RechargePackageUpdate,
|
||||
RechargePackageOut,
|
||||
)
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
router = APIRouter(tags=["recharge-packages"])
|
||||
|
||||
|
||||
def _to_out(pkg: RechargePackage) -> dict:
|
||||
return {
|
||||
"id": pkg.id,
|
||||
"name": pkg.name,
|
||||
"credits": round(pkg.credits, 2),
|
||||
"price": round(pkg.price, 2),
|
||||
"bonus_credits": round(pkg.bonus_credits, 2),
|
||||
"total_credits": round(pkg.credits + pkg.bonus_credits, 2),
|
||||
"description": pkg.description,
|
||||
"package_type": pkg.package_type,
|
||||
"is_gift": pkg.is_gift,
|
||||
"is_active": pkg.is_active,
|
||||
"sort_order": pkg.sort_order,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/recharge-packages")
|
||||
async def list_active_packages(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Public: list active recharge packages."""
|
||||
result = await db.execute(
|
||||
select(RechargePackage)
|
||||
.where(RechargePackage.is_active == True)
|
||||
.order_by(RechargePackage.sort_order)
|
||||
)
|
||||
return [_to_out(p) for p in result.scalars().all()]
|
||||
|
||||
|
||||
@router.get("/admin/recharge-packages")
|
||||
async def admin_list_packages(
|
||||
_admin=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Admin: list all packages."""
|
||||
result = await db.execute(
|
||||
select(RechargePackage).order_by(RechargePackage.sort_order)
|
||||
)
|
||||
return [_to_out(p) for p in result.scalars().all()]
|
||||
|
||||
|
||||
@router.post("/admin/recharge-packages")
|
||||
async def create_package(
|
||||
data: RechargePackageCreate,
|
||||
_admin=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
pkg = RechargePackage(id=generate_id(), **data.model_dump())
|
||||
db.add(pkg)
|
||||
await db.flush()
|
||||
return _to_out(pkg)
|
||||
|
||||
|
||||
@router.put("/admin/recharge-packages/{pkg_id}")
|
||||
async def update_package(
|
||||
pkg_id: str,
|
||||
data: RechargePackageUpdate,
|
||||
_admin=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(RechargePackage).where(RechargePackage.id == pkg_id)
|
||||
)
|
||||
pkg = result.scalar_one_or_none()
|
||||
if not pkg:
|
||||
raise HTTPException(status_code=404, detail="套餐不存在")
|
||||
for k, v in data.model_dump(exclude_unset=True).items():
|
||||
setattr(pkg, k, v)
|
||||
await db.flush()
|
||||
return _to_out(pkg)
|
||||
|
||||
|
||||
@router.delete("/admin/recharge-packages/{pkg_id}")
|
||||
async def delete_package(
|
||||
pkg_id: str,
|
||||
_admin=Depends(get_admin_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(RechargePackage).where(RechargePackage.id == pkg_id)
|
||||
)
|
||||
pkg = result.scalar_one_or_none()
|
||||
if not pkg:
|
||||
raise HTTPException(status_code=404, detail="套餐不存在")
|
||||
await db.delete(pkg)
|
||||
await db.flush()
|
||||
return {"ok": True}
|
||||
@@ -0,0 +1,45 @@
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
|
||||
from app.config import settings
|
||||
from app.schemas.sms import SmsSendRequest, SmsVerifyRequest, SmsResponse
|
||||
from app.services.sms import generate_and_send_sms, verify_sms_code
|
||||
|
||||
router = APIRouter(prefix="/sms", tags=["sms"])
|
||||
|
||||
|
||||
@router.post("/send", response_model=SmsResponse)
|
||||
async def send_sms_code(req: SmsSendRequest):
|
||||
"""Send SMS verification code. Requires captcha_token in production."""
|
||||
if not settings.SMS_MOCK:
|
||||
if not req.captcha_token:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="需要验证码",
|
||||
)
|
||||
from app.services.auth import decode_access_token
|
||||
token_sub = decode_access_token(req.captcha_token)
|
||||
if not token_sub or not token_sub.startswith("captcha:"):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="验证码无效",
|
||||
)
|
||||
|
||||
ok = await generate_and_send_sms(req.phone)
|
||||
if not ok:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="短信发送失败,请稍后重试",
|
||||
)
|
||||
return SmsResponse(message="验证码已发送", success=True)
|
||||
|
||||
|
||||
@router.post("/verify", response_model=SmsResponse)
|
||||
async def verify_sms(req: SmsVerifyRequest):
|
||||
"""Verify SMS code."""
|
||||
ok = await verify_sms_code(req.phone, req.code)
|
||||
if not ok:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="验证码错误或已过期",
|
||||
)
|
||||
return SmsResponse(message="验证成功", success=True)
|
||||
@@ -0,0 +1,53 @@
|
||||
import json
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from app.dependencies import get_db, get_current_user
|
||||
from app.models.user import User
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.schemas.video_engine import VideoEngineListResponse
|
||||
|
||||
router = APIRouter(prefix="/video-engines", tags=["video-engines"])
|
||||
|
||||
|
||||
@router.get("", response_model=VideoEngineListResponse)
|
||||
async def list_active_engines(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Public endpoint returning active video engine capabilities."""
|
||||
result = await db.execute(
|
||||
select(VideoEngine)
|
||||
.where(VideoEngine.is_active == True)
|
||||
.order_by(VideoEngine.priority.desc())
|
||||
)
|
||||
engines = result.scalars().all()
|
||||
items = []
|
||||
for e in engines:
|
||||
ratios = []
|
||||
resolutions = []
|
||||
durations = []
|
||||
try:
|
||||
ratios = json.loads(e.supported_ratios) if e.supported_ratios else []
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
resolutions = json.loads(e.supported_resolutions) if e.supported_resolutions else []
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
durations = json.loads(e.supported_durations) if e.supported_durations else []
|
||||
except Exception:
|
||||
pass
|
||||
items.append({
|
||||
"id": e.id,
|
||||
"name": e.name,
|
||||
"provider": e.provider,
|
||||
"supported_ratios": ratios,
|
||||
"supported_resolutions": resolutions,
|
||||
"supported_durations": durations,
|
||||
})
|
||||
return {"items": items}
|
||||
Reference in New Issue
Block a user