返回 MoneyPrinterTurbo
video.py
根目录 / app / controllers / v1 / video.py
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
401 lines PYTHON