| 1 | import glob |
| 2 | import os |
| 3 | import pathlib |
| 4 | import shutil |
| 5 | from typing import Union |
| 6 | |
| 7 | from fastapi import BackgroundTasks, Depends, Path, Query, Request, UploadFile |
| 8 | from fastapi.params import File |
| 9 | from fastapi.responses import FileResponse, StreamingResponse |
| 10 | from loguru import logger |
| 11 | |
| 12 | from app.config import config |
| 13 | from app.controllers import base |
| 14 | from app.controllers.manager.base_manager import TaskQueueFullError |
| 15 | from app.controllers.manager.memory_manager import InMemoryTaskManager |
| 16 | from app.controllers.manager.redis_manager import RedisTaskManager |
| 17 | from app.controllers.v1.base import new_router |
| 18 | from app.models.exception import HttpException |
| 19 | from app.models.schema import ( |
| 20 | AudioRequest, |
| 21 | BgmRetrieveResponse, |
| 22 | BgmUploadResponse, |
| 23 | SubtitleRequest, |
| 24 | TaskDeletionResponse, |
| 25 | TaskQueryRequest, |
| 26 | TaskQueryResponse, |
| 27 | TaskResponse, |
| 28 | TaskVideoRequest, |
| 29 | VideoMaterialUploadResponse, |
| 30 | VideoMaterialRetrieveResponse |
| 31 | ) |
| 32 | from app.services import state as sm |
| 33 | from app.services import task as tm |
| 34 | from app.utils import file_security, utils |
| 35 | |
| 36 | # 认证依赖项 |
| 37 | # router = new_router(dependencies=[Depends(base.verify_token)]) |
| 38 | router = new_router() |
| 39 | |
| 40 | _enable_redis = config.app.get("enable_redis", False) |
| 41 | _redis_host = config.app.get("redis_host", "localhost") |
| 42 | _redis_port = config.app.get("redis_port", 6379) |
| 43 | _redis_db = config.app.get("redis_db", 0) |
| 44 | _redis_password = config.app.get("redis_password", None) |
| 45 | _max_concurrent_tasks = config.app.get("max_concurrent_tasks", 5) |
| 46 | _max_queued_tasks = config.app.get("max_queued_tasks", 100) |
| 47 | |
| 48 | redis_url = f"redis://:{_redis_password}@{_redis_host}:{_redis_port}/{_redis_db}" |
| 49 | # 根据配置选择合适的任务管理器 |
| 50 | if _enable_redis: |
| 51 | task_manager = RedisTaskManager( |
| 52 | max_concurrent_tasks=_max_concurrent_tasks, |
| 53 | redis_url=redis_url, |
| 54 | max_queued_tasks=_max_queued_tasks, |
| 55 | ) |
| 56 | else: |
| 57 | task_manager = InMemoryTaskManager( |
| 58 | max_concurrent_tasks=_max_concurrent_tasks, |
| 59 | max_queued_tasks=_max_queued_tasks, |
| 60 | ) |
| 61 | |
| 62 | |
| 63 | def _sanitize_upload_filename(filename: str, request_id: str) -> str: |
| 64 | # 浏览器或客户端有时会附带目录信息,甚至可能夹带 ../ 这类穿越片段。 |
| 65 | # 这里只保留纯文件名,避免上传接口把文件写到目标目录之外。 |
| 66 | normalized_name = (filename or "").replace("\\", "/").split("/")[-1].strip() |
| 67 | if not normalized_name or normalized_name in {".", ".."}: |
| 68 | raise HttpException( |
| 69 | task_id=request_id, |
| 70 | status_code=400, |
| 71 | message=f"{request_id}: invalid filename", |
| 72 | ) |
| 73 | return normalized_name |
| 74 | |
| 75 | |
| 76 | def _resolve_path_within_directory(base_dir: str, unsafe_path: str, request_id: str) -> str: |
| 77 | try: |
| 78 | return file_security.resolve_path_within_directory(base_dir, unsafe_path) |
| 79 | except ValueError as exc: |
| 80 | logger.warning( |
| 81 | f"reject unsafe file path, request_id: {request_id}, path: {unsafe_path}, " |
| 82 | f"error: {str(exc)}" |
| 83 | ) |
| 84 | raise HttpException( |
| 85 | task_id=request_id, |
| 86 | status_code=404 if str(exc) == "file does not exist" else 403, |
| 87 | message=f"{request_id}: invalid file path", |
| 88 | ) |
| 89 | |
| 90 | def _task_file_to_uri(file: str, endpoint: str, task_dir: str, request_id: str) -> str: |
| 91 | if not isinstance(file, str): |
| 92 | return file |
| 93 | |
| 94 | if file.startswith(("http://", "https://")): |
| 95 | return file |
| 96 | |
| 97 | try: |
| 98 | resolved_path = file_security.resolve_path_within_directory(task_dir, file) |
| 99 | except ValueError as exc: |
| 100 | # 任务状态理论上只应保存任务目录内的产物路径。这里不再继续拼接 URL, |
| 101 | # 避免把异常路径包装成可访问链接;同时保留原值,便于排查历史脏数据。 |
| 102 | logger.warning( |
| 103 | f"skip unsafe task output path, request_id: {request_id}, path: {file}, " |
| 104 | f"error: {str(exc)}" |
| 105 | ) |
| 106 | return file |
| 107 | |
| 108 | relative_path = os.path.relpath(resolved_path, task_dir).replace("\\", "/") |
| 109 | uri_path = f"tasks/{relative_path}" |
| 110 | if endpoint: |
| 111 | return f"{endpoint.rstrip('/')}/{uri_path}" |
| 112 | return f"/{uri_path}" |
| 113 | |
| 114 | |
| 115 | @router.post("/videos", response_model=TaskResponse, summary="Generate a short video") |
| 116 | def create_video( |
| 117 | background_tasks: BackgroundTasks, request: Request, body: TaskVideoRequest |
| 118 | ): |
| 119 | return create_task(request, body, stop_at="video") |
| 120 | |
| 121 | |
| 122 | @router.post("/subtitle", response_model=TaskResponse, summary="Generate subtitle only") |
| 123 | def create_subtitle( |
| 124 | background_tasks: BackgroundTasks, request: Request, body: SubtitleRequest |
| 125 | ): |
| 126 | return create_task(request, body, stop_at="subtitle") |
| 127 | |
| 128 | |
| 129 | @router.post("/audio", response_model=TaskResponse, summary="Generate audio only") |
| 130 | def create_audio( |
| 131 | background_tasks: BackgroundTasks, request: Request, body: AudioRequest |
| 132 | ): |
| 133 | return create_task(request, body, stop_at="audio") |
| 134 | |
| 135 | |
| 136 | def create_task( |
| 137 | request: Request, |
| 138 | body: Union[TaskVideoRequest, SubtitleRequest, AudioRequest], |
| 139 | stop_at: str, |
| 140 | ): |
| 141 | task_id = utils.get_uuid() |
| 142 | request_id = base.get_task_id(request) |
| 143 | try: |
| 144 | task = { |
| 145 | "task_id": task_id, |
| 146 | "request_id": request_id, |
| 147 | "params": body.model_dump(), |
| 148 | } |
| 149 | sm.state.update_task(task_id) |
| 150 | task_manager.add_task(tm.start, task_id=task_id, params=body, stop_at=stop_at) |
| 151 | logger.success(f"Task created: {utils.to_json(task)}") |
| 152 | return utils.get_response(200, task) |
| 153 | except TaskQueueFullError as e: |
| 154 | sm.state.delete_task(task_id) |
| 155 | logger.warning( |
| 156 | f"reject task because queue is full, request_id: {request_id}, task_id: {task_id}" |
| 157 | ) |
| 158 | raise HttpException( |
| 159 | task_id=task_id, status_code=429, message=f"{request_id}: {str(e)}" |
| 160 | ) |
| 161 | except ValueError as e: |
| 162 | raise HttpException( |
| 163 | task_id=task_id, status_code=400, message=f"{request_id}: {str(e)}" |
| 164 | ) |
| 165 | |
| 166 | @router.get("/tasks", response_model=TaskQueryResponse, summary="Get all tasks") |
| 167 | def get_all_tasks(request: Request, page: int = Query(1, ge=1), page_size: int = Query(10, ge=1)): |
| 168 | tasks, total = sm.state.get_all_tasks(page, page_size) |
| 169 | |
| 170 | response = { |
| 171 | "tasks": tasks, |
| 172 | "total": total, |
| 173 | "page": page, |
| 174 | "page_size": page_size, |
| 175 | } |
| 176 | return utils.get_response(200, response) |
| 177 | |
| 178 | |
| 179 | |
| 180 | @router.get( |
| 181 | "/tasks/{task_id}", response_model=TaskQueryResponse, summary="Query task status" |
| 182 | ) |
| 183 | def get_task( |
| 184 | request: Request, |
| 185 | task_id: str = Path(..., description="Task ID"), |
| 186 | query: TaskQueryRequest = Depends(), |
| 187 | ): |
| 188 | request_id = base.get_task_id(request) |
| 189 | endpoint = config.app.get("endpoint", "").rstrip("/") |
| 190 | task = sm.state.get_task(task_id) |
| 191 | if task: |
| 192 | task_dir = utils.task_dir() |
| 193 | response_task = dict(task) |
| 194 | |
| 195 | if "videos" in task: |
| 196 | response_task["videos"] = [ |
| 197 | _task_file_to_uri(v, endpoint, task_dir, request_id) |
| 198 | for v in task["videos"] |
| 199 | ] |
| 200 | if "combined_videos" in task: |
| 201 | response_task["combined_videos"] = [ |
| 202 | _task_file_to_uri(v, endpoint, task_dir, request_id) |
| 203 | for v in task["combined_videos"] |
| 204 | ] |
| 205 | return utils.get_response(200, response_task) |
| 206 | |
| 207 | raise HttpException( |
| 208 | task_id=task_id, status_code=404, message=f"{request_id}: task not found" |
| 209 | ) |
| 210 | |
| 211 | |
| 212 | @router.delete( |
| 213 | "/tasks/{task_id}", |
| 214 | response_model=TaskDeletionResponse, |
| 215 | summary="Delete a generated short video task", |
| 216 | ) |
| 217 | def delete_video(request: Request, task_id: str = Path(..., description="Task ID")): |
| 218 | request_id = base.get_task_id(request) |
| 219 | task = sm.state.get_task(task_id) |
| 220 | if task: |
| 221 | tasks_dir = utils.task_dir() |
| 222 | current_task_dir = os.path.join(tasks_dir, task_id) |
| 223 | if os.path.exists(current_task_dir): |
| 224 | shutil.rmtree(current_task_dir) |
| 225 | |
| 226 | sm.state.delete_task(task_id) |
| 227 | logger.success(f"video deleted: {utils.to_json(task)}") |
| 228 | return utils.get_response(200) |
| 229 | |
| 230 | raise HttpException( |
| 231 | task_id=task_id, status_code=404, message=f"{request_id}: task not found" |
| 232 | ) |
| 233 | |
| 234 | |
| 235 | @router.get( |
| 236 | "/musics", response_model=BgmRetrieveResponse, summary="Retrieve local BGM files" |
| 237 | ) |
| 238 | def get_bgm_list(request: Request): |
| 239 | suffix = "*.mp3" |
| 240 | song_dir = utils.song_dir() |
| 241 | files = glob.glob(os.path.join(song_dir, suffix)) |
| 242 | bgm_list = [] |
| 243 | for file in files: |
| 244 | filename = os.path.basename(file) |
| 245 | bgm_list.append( |
| 246 | { |
| 247 | "name": filename, |
| 248 | "size": os.path.getsize(file), |
| 249 | # 只返回文件名,避免把服务器绝对路径暴露给调用方。 |
| 250 | # 服务端后续会把该文件名解析回 songs 白名单目录。 |
| 251 | "file": filename, |
| 252 | } |
| 253 | ) |
| 254 | response = {"files": bgm_list} |
| 255 | return utils.get_response(200, response) |
| 256 | |
| 257 | |
| 258 | @router.post( |
| 259 | "/musics", |
| 260 | response_model=BgmUploadResponse, |
| 261 | summary="Upload the BGM file to the songs directory", |
| 262 | ) |
| 263 | def upload_bgm_file(request: Request, file: UploadFile = File(...)): |
| 264 | request_id = base.get_task_id(request) |
| 265 | safe_filename = _sanitize_upload_filename(file.filename, request_id) |
| 266 | # check file ext |
| 267 | if safe_filename.lower().endswith("mp3"): |
| 268 | song_dir = utils.song_dir() |
| 269 | save_path = os.path.join(song_dir, safe_filename) |
| 270 | # save file |
| 271 | with open(save_path, "wb+") as buffer: |
| 272 | # If the file already exists, it will be overwritten |
| 273 | file.file.seek(0) |
| 274 | buffer.write(file.file.read()) |
| 275 | response = {"file": safe_filename} |
| 276 | return utils.get_response(200, response) |
| 277 | |
| 278 | raise HttpException( |
| 279 | "", status_code=400, message=f"{request_id}: Only *.mp3 files can be uploaded" |
| 280 | ) |
| 281 | |
| 282 | @router.get( |
| 283 | "/video_materials", response_model=VideoMaterialRetrieveResponse, summary="Retrieve local video materials" |
| 284 | ) |
| 285 | def get_video_materials_list(request: Request): |
| 286 | allowed_suffixes = ("mp4", "mov", "avi", "flv", "mkv", "jpg", "jpeg", "png") |
| 287 | local_videos_dir = utils.storage_dir("local_videos", create=True) |
| 288 | files = [] |
| 289 | for suffix in allowed_suffixes: |
| 290 | files.extend(glob.glob(os.path.join(local_videos_dir, f"*.{suffix}"))) |
| 291 | # 文件系统枚举顺序不稳定,直接返回会导致“顺序拼接”在不同机器或不同 |
| 292 | # 时刻表现不一致。这里统一按文件名排序,至少保证服务端返回顺序可预测。 |
| 293 | files.sort(key=lambda file_path: os.path.basename(file_path).lower()) |
| 294 | video_materials_list = [] |
| 295 | for file in files: |
| 296 | filename = os.path.basename(file) |
| 297 | video_materials_list.append( |
| 298 | { |
| 299 | "name": filename, |
| 300 | "size": os.path.getsize(file), |
| 301 | # 与 BGM 一样,只返回文件名;创建任务时再在 local_videos |
| 302 | # 白名单目录内解析,避免 API 泄露宿主机绝对路径。 |
| 303 | "file": filename, |
| 304 | } |
| 305 | ) |
| 306 | response = {"files": video_materials_list} |
| 307 | return utils.get_response(200, response) |
| 308 | |
| 309 | |
| 310 | @router.post( |
| 311 | "/video_materials", |
| 312 | response_model=VideoMaterialUploadResponse, |
| 313 | summary="Upload the video material file to the local videos directory", |
| 314 | ) |
| 315 | def upload_video_material_file(request: Request, file: UploadFile = File(...)): |
| 316 | request_id = base.get_task_id(request) |
| 317 | safe_filename = _sanitize_upload_filename(file.filename, request_id) |
| 318 | # check file ext |
| 319 | allowed_suffixes = ("mp4", "mov", "avi", "flv", "mkv", "jpg", "jpeg", "png") |
| 320 | normalized_filename = safe_filename.lower() |
| 321 | # 统一按小写扩展名校验,兼容 .MOV 这类大写后缀文件。 |
| 322 | if normalized_filename.endswith(allowed_suffixes): |
| 323 | local_videos_dir = utils.storage_dir("local_videos", create=True) |
| 324 | save_path = os.path.join(local_videos_dir, safe_filename) |
| 325 | # save file |
| 326 | with open(save_path, "wb+") as buffer: |
| 327 | # If the file already exists, it will be overwritten |
| 328 | file.file.seek(0) |
| 329 | buffer.write(file.file.read()) |
| 330 | response = {"file": safe_filename} |
| 331 | return utils.get_response(200, response) |
| 332 | |
| 333 | raise HttpException( |
| 334 | "", status_code=400, message=f"{request_id}: Only files with extensions {', '.join(allowed_suffixes)} can be uploaded" |
| 335 | ) |
| 336 | |
| 337 | @router.get("/stream/{file_path:path}") |
| 338 | async def stream_video(request: Request, file_path: str): |
| 339 | request_id = base.get_task_id(request) |
| 340 | tasks_dir = utils.task_dir() |
| 341 | video_path = _resolve_path_within_directory(tasks_dir, file_path, request_id) |
| 342 | range_header = request.headers.get("Range") |
| 343 | video_size = os.path.getsize(video_path) |
| 344 | start, end = 0, video_size - 1 |
| 345 | |
| 346 | length = video_size |
| 347 | if range_header: |
| 348 | range_ = range_header.split("bytes=")[1] |
| 349 | start, end = [int(part) if part else None for part in range_.split("-")] |
| 350 | if start is None: |
| 351 | start = video_size - end |
| 352 | end = video_size - 1 |
| 353 | if end is None: |
| 354 | end = video_size - 1 |
| 355 | length = end - start + 1 |
| 356 | |
| 357 | def file_iterator(file_path, offset=0, bytes_to_read=None): |
| 358 | with open(file_path, "rb") as f: |
| 359 | f.seek(offset, os.SEEK_SET) |
| 360 | remaining = bytes_to_read or video_size |
| 361 | while remaining > 0: |
| 362 | bytes_to_read = min(4096, remaining) |
| 363 | data = f.read(bytes_to_read) |
| 364 | if not data: |
| 365 | break |
| 366 | remaining -= len(data) |
| 367 | yield data |
| 368 | |
| 369 | response = StreamingResponse( |
| 370 | file_iterator(video_path, start, length), media_type="video/mp4" |
| 371 | ) |
| 372 | response.headers["Content-Range"] = f"bytes {start}-{end}/{video_size}" |
| 373 | response.headers["Accept-Ranges"] = "bytes" |
| 374 | response.headers["Content-Length"] = str(length) |
| 375 | response.status_code = 206 # Partial Content |
| 376 | |
| 377 | return response |
| 378 | |
| 379 | |
| 380 | @router.get("/download/{file_path:path}") |
| 381 | async def download_video(request: Request, file_path: str): |
| 382 | """ |
| 383 | download video |
| 384 | :param request: Request request |
| 385 | :param file_path: video file path, eg: /cd1727ed-3473-42a2-a7da-4faafafec72b/final-1.mp4 |
| 386 | :return: video file |
| 387 | """ |
| 388 | request_id = base.get_task_id(request) |
| 389 | tasks_dir = utils.task_dir() |
| 390 | video_path = _resolve_path_within_directory(tasks_dir, file_path, request_id) |
| 391 | file_path = pathlib.Path(video_path) |
| 392 | filename = file_path.stem |
| 393 | extension = file_path.suffix |
| 394 | headers = {"Content-Disposition": f"attachment; filename={filename}{extension}"} |
| 395 | return FileResponse( |
| 396 | path=video_path, |
| 397 | headers=headers, |
| 398 | filename=f"{filename}{extension}", |
| 399 | media_type=f"video/{extension[1:]}", |
| 400 | ) |
| 401 |