Replace the prior release tree with the current plugin, frontend, tests, and documentation. Document retired node IDs and the public Gitea update source.
510 lines
18 KiB
Python
510 lines
18 KiB
Python
"""
|
|
Veo 视频生成 API 客户端
|
|
提供视频创建、状态轮询、视频下载功能
|
|
"""
|
|
|
|
import asyncio
|
|
import base64
|
|
import json
|
|
import os
|
|
import time
|
|
from typing import Any, Callable, Dict, List, Optional
|
|
|
|
import aiohttp
|
|
|
|
from .base_client import BaseAPIClient
|
|
from ..utils.config import get_api_key_or_raise, get_api_base_url
|
|
from ..utils.image_utils import encode_image_to_base64
|
|
from ..utils.video_task import PollDeadline, download_video_to_file
|
|
|
|
|
|
class VeoClient(BaseAPIClient):
|
|
"""
|
|
Veo 视频生成客户端
|
|
|
|
工作流程:
|
|
1. create_video → POST /v1/videos (提交生成任务)
|
|
2. poll_status → GET /v1/videos/{id} (轮询直到完成/失败)
|
|
3. download_video→ GET /v1/videos/{id}/content (下载视频文件)
|
|
"""
|
|
|
|
CREATE_ENDPOINT = "/v1/videos"
|
|
STATUS_ENDPOINT = "/v1/videos/{video_id}"
|
|
CONTENT_ENDPOINT = "/v1/videos/{video_id}/content"
|
|
|
|
POLL_INITIAL_INTERVAL = 3
|
|
POLL_MAX_INTERVAL = 15
|
|
|
|
def __init__(self):
|
|
api_key = get_api_key_or_raise()
|
|
base_url = get_api_base_url()
|
|
super().__init__(base_url=base_url, api_key=api_key)
|
|
|
|
# ------------------------------------------------------------------
|
|
# BaseAPIClient 抽象方法实现
|
|
# ------------------------------------------------------------------
|
|
|
|
def get_endpoint(self, **kwargs) -> str:
|
|
return self.CREATE_ENDPOINT
|
|
|
|
def build_request_body(self, **kwargs) -> Dict[str, Any]:
|
|
return {}
|
|
|
|
def parse_response(self, response: Dict[str, Any]) -> Any:
|
|
return response
|
|
|
|
# ------------------------------------------------------------------
|
|
# 核心异步方法
|
|
# ------------------------------------------------------------------
|
|
|
|
async def create_video_async(
|
|
self,
|
|
prompt: str,
|
|
model: str,
|
|
seconds: int = 8,
|
|
size: str = "720x1280",
|
|
first_frame_bytes: Optional[bytes] = None,
|
|
last_frame_bytes: Optional[bytes] = None,
|
|
reference_bytes: Optional[bytes] = None,
|
|
seed: Optional[int] = None,
|
|
session: Optional[aiohttp.ClientSession] = None,
|
|
) -> Dict[str, Any]:
|
|
"""
|
|
提交视频生成任务
|
|
|
|
格式策略:
|
|
- 无参考图片:application/json
|
|
- 有参考图片:multipart/form-data,图片以 PNG 文件上传
|
|
|
|
Args:
|
|
prompt: 提示词
|
|
model: 模型名称
|
|
seconds: 视频时长(秒)
|
|
size: 分辨率
|
|
first_frame_bytes: 首帧图片字节
|
|
last_frame_bytes: 尾帧图片字节
|
|
reference_bytes: 参考图片字节
|
|
seed: 随机种子
|
|
session: aiohttp 会话
|
|
|
|
Returns:
|
|
API 响应 JSON,包含 video id 和初始状态
|
|
"""
|
|
url = f"{self.base_url}{self.CREATE_ENDPOINT}"
|
|
headers = {"Authorization": f"Bearer {self.api_key}"}
|
|
|
|
# 检查是否有图片
|
|
has_images = any([first_frame_bytes, last_frame_bytes, reference_bytes])
|
|
|
|
if has_images:
|
|
# 有图片:multipart/form-data + PNG 文件上传
|
|
if first_frame_bytes and len(first_frame_bytes) > self.max_request_size:
|
|
raise ValueError(f"首帧图片过大,超过 {self.max_request_size / 1024 / 1024:.0f}MB 限制")
|
|
if last_frame_bytes and len(last_frame_bytes) > self.max_request_size:
|
|
raise ValueError(f"尾帧图片过大,超过 {self.max_request_size / 1024 / 1024:.0f}MB 限制")
|
|
if reference_bytes and len(reference_bytes) > self.max_request_size:
|
|
raise ValueError(f"参考图片过大,超过 {self.max_request_size / 1024 / 1024:.0f}MB 限制")
|
|
|
|
form = aiohttp.FormData()
|
|
form.add_field("prompt", prompt)
|
|
form.add_field("model", model)
|
|
form.add_field("seconds", str(seconds))
|
|
form.add_field("size", size)
|
|
# 注意:seed 不被上游 API 接受,仅在 ComfyUI 节点侧用于缓存刷新
|
|
# if seed is not None:
|
|
# form.add_field("seed", str(seed))
|
|
|
|
# 使用 input_reference 字段(OpenAI兼容格式)
|
|
# 尝试支持多张图片:按顺序添加多个 input_reference 字段
|
|
if first_frame_bytes:
|
|
form.add_field(
|
|
"input_reference",
|
|
first_frame_bytes,
|
|
filename="first_frame.png",
|
|
content_type="image/png",
|
|
)
|
|
if last_frame_bytes:
|
|
form.add_field(
|
|
"input_reference",
|
|
last_frame_bytes,
|
|
filename="last_frame.png",
|
|
content_type="image/png",
|
|
)
|
|
if reference_bytes:
|
|
form.add_field(
|
|
"input_reference",
|
|
reference_bytes,
|
|
filename="reference.png",
|
|
content_type="image/png",
|
|
)
|
|
|
|
send_kwargs: Dict[str, Any] = {"data": form, "headers": headers}
|
|
else:
|
|
# 无图片:application/json
|
|
body: Dict[str, Any] = {
|
|
"model": model,
|
|
"prompt": prompt,
|
|
"seconds": str(seconds),
|
|
"size": size,
|
|
}
|
|
# 注意:seed 不被上游 API 接受,仅在 ComfyUI 节点侧用于缓存刷新
|
|
# if seed is not None:
|
|
# body["seed"] = str(seed)
|
|
send_kwargs = {"json": body, "headers": headers}
|
|
|
|
# 打印请求调试信息
|
|
import json
|
|
if has_images:
|
|
print(f"Veo: 使用 multipart/form-data 格式上传图片")
|
|
else:
|
|
print(f"Veo API 请求体: {json.dumps(body, ensure_ascii=False)}")
|
|
|
|
close_session = False
|
|
if session is None:
|
|
session = self._make_session()
|
|
close_session = True
|
|
|
|
try:
|
|
async with session.post(url, **send_kwargs) as response:
|
|
if response.status != 200:
|
|
error_text = await response.text()
|
|
error_message = self._extract_error_message(error_text, response.status)
|
|
raise RuntimeError(error_message)
|
|
|
|
resp_json = await response.json()
|
|
return resp_json
|
|
|
|
finally:
|
|
if close_session:
|
|
await session.close()
|
|
|
|
async def poll_video_status_async(
|
|
self,
|
|
video_id: str,
|
|
progress_callback: Optional[Callable[[int, float], None]] = None,
|
|
session: Optional[aiohttp.ClientSession] = None,
|
|
) -> Dict[str, Any]:
|
|
"""
|
|
轮询视频生成状态,直到完成或失败
|
|
|
|
Args:
|
|
video_id: 视频任务 ID
|
|
progress_callback: 进度回调 (progress_percent, elapsed_seconds)
|
|
session: aiohttp 会话
|
|
|
|
Returns:
|
|
最终状态的 API 响应
|
|
|
|
Raises:
|
|
RuntimeError: 生成失败
|
|
"""
|
|
url = f"{self.base_url}{self.STATUS_ENDPOINT.format(video_id=video_id)}"
|
|
headers = self.get_headers(use_bearer_token=True)
|
|
|
|
close_session = False
|
|
if session is None:
|
|
session = self._make_session()
|
|
close_session = True
|
|
|
|
interval = self.POLL_INITIAL_INTERVAL
|
|
deadline = PollDeadline(label="Veo 视频")
|
|
|
|
try:
|
|
while True:
|
|
deadline.check()
|
|
async with session.get(url, headers=headers) as response:
|
|
if response.status != 200:
|
|
error_text = await response.text()
|
|
error_message = self._extract_error_message(error_text, response.status)
|
|
raise RuntimeError(error_message)
|
|
|
|
data = await response.json()
|
|
|
|
# status 兼容大小写
|
|
status = data.get("status", "").lower()
|
|
|
|
# progress 兼容整数和字符串
|
|
progress_raw = data.get("progress", 0)
|
|
if isinstance(progress_raw, str):
|
|
try:
|
|
progress = int(progress_raw.rstrip("%").strip())
|
|
except ValueError:
|
|
progress = 0
|
|
else:
|
|
progress = int(progress_raw) if progress_raw else 0
|
|
|
|
if progress_callback:
|
|
progress_callback(progress)
|
|
|
|
if status == "completed":
|
|
return data
|
|
|
|
if status == "failed":
|
|
error_info = data.get("error", {})
|
|
error_msg = error_info.get("message", "未知错误") if isinstance(error_info, dict) else str(error_info)
|
|
raise RuntimeError(f"视频生成失败: {error_msg}")
|
|
|
|
await asyncio.sleep(interval)
|
|
interval = min(interval * 1.5, self.POLL_MAX_INTERVAL)
|
|
|
|
finally:
|
|
if close_session:
|
|
await session.close()
|
|
|
|
async def download_video_async(
|
|
self,
|
|
video_id: str,
|
|
save_path: str,
|
|
session: Optional[aiohttp.ClientSession] = None,
|
|
) -> str:
|
|
"""
|
|
下载生成的视频文件
|
|
|
|
Returns:
|
|
保存的文件路径
|
|
"""
|
|
url = f"{self.base_url}{self.CONTENT_ENDPOINT.format(video_id=video_id)}"
|
|
headers = self.get_headers(use_bearer_token=True)
|
|
|
|
close_session = False
|
|
if session is None:
|
|
session = self._make_session()
|
|
close_session = True
|
|
|
|
try:
|
|
async with session.get(url, headers=headers, allow_redirects=True) as response:
|
|
if response.status != 200:
|
|
error_text = await response.text()
|
|
error_message = self._extract_error_message(error_text, response.status)
|
|
raise RuntimeError(f"视频下载失败: {error_message}")
|
|
|
|
content_type = response.headers.get("Content-Type", "")
|
|
download_url = None
|
|
if "application/json" in content_type:
|
|
data = await response.json()
|
|
download_url = data.get("url") or data.get("download_url")
|
|
if not download_url:
|
|
raise RuntimeError("视频下载失败: 响应中未找到下载链接")
|
|
|
|
if download_url:
|
|
await self._download_from_url(download_url, save_path, session)
|
|
else:
|
|
# content 端点直接返回视频流(幂等 GET,可安全重连续传)
|
|
await download_video_to_file(
|
|
session, url, save_path, headers=headers, label="VEO 视频",
|
|
)
|
|
|
|
return save_path
|
|
|
|
finally:
|
|
if close_session:
|
|
await session.close()
|
|
|
|
# ------------------------------------------------------------------
|
|
# 同步包装
|
|
# ------------------------------------------------------------------
|
|
|
|
def generate_video_sync(
|
|
self,
|
|
prompt: str,
|
|
model: str,
|
|
seconds: int,
|
|
size: str,
|
|
save_path: str,
|
|
first_frame_bytes: Optional[bytes] = None,
|
|
last_frame_bytes: Optional[bytes] = None,
|
|
reference_bytes: Optional[bytes] = None,
|
|
seed: Optional[int] = None,
|
|
progress_callback: Optional[Callable[[int], None]] = None,
|
|
on_stage: Optional[Callable[[str], None]] = None,
|
|
) -> str:
|
|
"""
|
|
同步执行完整的视频生成流程(创建 → 轮询 → 下载)
|
|
"""
|
|
|
|
async def _run():
|
|
connector = aiohttp.TCPConnector(ssl=False, limit=0)
|
|
async with aiohttp.ClientSession(connector=connector) as session:
|
|
# 1. 提交任务
|
|
if on_stage:
|
|
on_stage("submitting")
|
|
result = await self.create_video_async(
|
|
prompt=prompt,
|
|
model=model,
|
|
seconds=seconds,
|
|
size=size,
|
|
first_frame_bytes=first_frame_bytes,
|
|
last_frame_bytes=last_frame_bytes,
|
|
reference_bytes=reference_bytes,
|
|
seed=seed,
|
|
session=session,
|
|
)
|
|
video_id = result.get("id")
|
|
if not video_id:
|
|
raise RuntimeError("API 未返回视频任务 ID")
|
|
|
|
if on_stage:
|
|
on_stage(f"submitted:{video_id}")
|
|
|
|
# 2. 轮询状态
|
|
if on_stage:
|
|
on_stage("polling")
|
|
await self.poll_video_status_async(
|
|
video_id=video_id,
|
|
progress_callback=progress_callback,
|
|
session=session,
|
|
)
|
|
|
|
# 3. 下载视频
|
|
if on_stage:
|
|
on_stage("downloading")
|
|
path = await self.download_video_async(
|
|
video_id=video_id,
|
|
save_path=save_path,
|
|
session=session,
|
|
)
|
|
|
|
if on_stage:
|
|
on_stage("done")
|
|
return path
|
|
|
|
return self.run_async_in_thread(_run())
|
|
|
|
def generate_batch_videos_sync(
|
|
self,
|
|
prompt: str,
|
|
model: str,
|
|
seconds: int,
|
|
size: str,
|
|
save_paths: List[str],
|
|
first_frame_bytes: Optional[bytes] = None,
|
|
last_frame_bytes: Optional[bytes] = None,
|
|
reference_bytes: Optional[bytes] = None,
|
|
seed: Optional[int] = None,
|
|
progress_callback: Optional[Callable[[int, int, bool, Optional[str]], None]] = None,
|
|
) -> List[str]:
|
|
"""
|
|
同步并发生成多个视频
|
|
"""
|
|
async def _run():
|
|
batch_size = len(save_paths)
|
|
connector = aiohttp.TCPConnector(ssl=False, limit=0)
|
|
|
|
async def generate_one(save_path: str):
|
|
return await self._generate_one_video_async(
|
|
prompt=prompt,
|
|
model=model,
|
|
seconds=seconds,
|
|
size=size,
|
|
save_path=save_path,
|
|
first_frame_bytes=first_frame_bytes,
|
|
last_frame_bytes=last_frame_bytes,
|
|
reference_bytes=reference_bytes,
|
|
seed=seed,
|
|
)
|
|
|
|
async with aiohttp.ClientSession(connector=connector) as session:
|
|
tasks = [generate_one(p) for p in save_paths]
|
|
results = await asyncio.gather(*tasks, return_exceptions=True)
|
|
|
|
completed = 0
|
|
paths: List[str] = []
|
|
first_error = None
|
|
for i, result in enumerate(results):
|
|
if isinstance(result, Exception):
|
|
error_msg = str(result)
|
|
print(f"Veo: 第 {i + 1} 个视频生成失败")
|
|
if first_error is None:
|
|
first_error = result
|
|
if progress_callback:
|
|
progress_callback(i + 1, batch_size, False, error_msg)
|
|
else:
|
|
completed += 1
|
|
paths.append(result)
|
|
if progress_callback:
|
|
progress_callback(completed, batch_size, True, None)
|
|
|
|
if not paths:
|
|
if first_error:
|
|
raise first_error
|
|
raise RuntimeError(f"批量视频生成失败,{batch_size} 个任务全部失败")
|
|
|
|
return paths
|
|
|
|
return self.run_async_in_thread(_run())
|
|
|
|
async def _generate_one_video_async(
|
|
self,
|
|
prompt: str,
|
|
model: str,
|
|
seconds: int,
|
|
size: str,
|
|
save_path: str,
|
|
first_frame_bytes: Optional[bytes] = None,
|
|
last_frame_bytes: Optional[bytes] = None,
|
|
reference_bytes: Optional[bytes] = None,
|
|
seed: Optional[int] = None,
|
|
session: Optional[aiohttp.ClientSession] = None,
|
|
) -> str:
|
|
"""异步生成单个视频"""
|
|
result = await self.create_video_async(
|
|
prompt=prompt,
|
|
model=model,
|
|
seconds=seconds,
|
|
size=size,
|
|
first_frame_bytes=first_frame_bytes,
|
|
last_frame_bytes=last_frame_bytes,
|
|
reference_bytes=reference_bytes,
|
|
seed=seed,
|
|
session=session,
|
|
)
|
|
video_id = result.get("id")
|
|
if not video_id:
|
|
raise RuntimeError("API 未返回视频任务 ID")
|
|
|
|
await self.poll_video_status_async(video_id=video_id, session=session)
|
|
path = await self.download_video_async(
|
|
video_id=video_id, save_path=save_path, session=session
|
|
)
|
|
return path
|
|
|
|
# ------------------------------------------------------------------
|
|
# 内部辅助方法
|
|
# ------------------------------------------------------------------
|
|
|
|
async def _download_from_url(
|
|
self,
|
|
url: str,
|
|
save_path: str,
|
|
session: aiohttp.ClientSession,
|
|
) -> None:
|
|
"""从给定 URL 下载文件到本地路径"""
|
|
await download_video_to_file(session, url, save_path, label="VEO 视频")
|
|
|
|
@staticmethod
|
|
def _extract_error_message(error_text: str, status_code: int) -> str:
|
|
"""从错误响应中提取可读的错误信息"""
|
|
error_message = error_text
|
|
try:
|
|
error_json = json.loads(error_text)
|
|
if "error" in error_json:
|
|
if isinstance(error_json["error"], dict):
|
|
error_message = error_json["error"].get("message", error_text)
|
|
else:
|
|
error_message = str(error_json["error"])
|
|
elif "message" in error_json:
|
|
error_message = error_json["message"]
|
|
except (json.JSONDecodeError, KeyError):
|
|
pass
|
|
|
|
status_hints = {
|
|
400: "请求参数错误 (400)",
|
|
401: "认证失败 (401),请检查 API 密钥",
|
|
403: "权限不足 (403),请检查账户权限或余额",
|
|
429: "请求频率超限 (429),请稍后重试",
|
|
503: "服务暂时不可用 (503),请稍后重试",
|
|
504: "请求超时 (504),请稍后重试",
|
|
}
|
|
hint = status_hints.get(status_code, f"API 请求失败 (状态码: {status_code})")
|
|
return f"{hint}\nAPI 返回: {error_message}"
|