This commit is contained in:
2026-05-25 17:08:18 +08:00
parent df501f6151
commit f1259b71b5
9178 changed files with 1626125 additions and 0 deletions
View File
+32
View File
@@ -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
+204
View File
@@ -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)}
+20
View File
@@ -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)
+41
View File
@@ -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
+602
View File
@@ -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"}
+48
View File
@@ -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}
+46
View File
@@ -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
+92
View File
@@ -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"}
+142
View File
@@ -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)
+81
View File
@@ -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()
+68
View File
@@ -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}
+45
View File
@@ -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)
+53
View File
@@ -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}