- 新增启动弹窗通知(绿色主题,支持关闭) - 新增 StreamPreview 流式文本预览节点 - 新增 fileUpload、updateNotifier 前端 JS 模块 - 重构多个 client,统一错误处理 - 删除废弃节点 batch_nano_banana_v2、quan_neng_sheng_tu 等 - 将 .config 纳入版本控制(已清空密钥)
511 lines
18 KiB
Python
511 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
|
|
|
|
|
|
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
|
|
|
|
try:
|
|
while True:
|
|
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", "")
|
|
|
|
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("视频下载失败: 响应中未找到下载链接")
|
|
await self._download_from_url(download_url, save_path, session)
|
|
else:
|
|
os.makedirs(os.path.dirname(save_path), exist_ok=True)
|
|
with open(save_path, "wb") as f:
|
|
async for chunk in response.content.iter_chunked(8192):
|
|
f.write(chunk)
|
|
|
|
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 下载文件到本地路径"""
|
|
os.makedirs(os.path.dirname(save_path), exist_ok=True)
|
|
async with session.get(url) as response:
|
|
if response.status != 200:
|
|
raise RuntimeError(f"从下载链接获取视频失败 (状态码: {response.status})")
|
|
with open(save_path, "wb") as f:
|
|
async for chunk in response.content.iter_chunked(8192):
|
|
f.write(chunk)
|
|
|
|
@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}"
|