Skip to content

Commit 794ba91

Browse files
committed
feat: enhance GenerationVideoModel with API version detection and session management
1 parent b58a6db commit 794ba91

1 file changed

Lines changed: 175 additions & 92 deletions

File tree

  • apps/models_provider/impl/minimax_model_provider/model
Lines changed: 175 additions & 92 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
1+
# coding=utf-8
12
import time
2-
from typing import Dict
3+
from typing import Dict, Optional
34

45
import requests
56

@@ -9,21 +10,36 @@
910

1011

1112
class GenerationVideoModel(MaxKBBaseModel, BaseGenerationVideo):
13+
"""MiniMax 视频生成模型,兼容 V1 与 V2 (MiniMax-H3) 两套接口。"""
14+
1215
api_key: str
1316
api_base: str
1417
model_name: str
1518
params: dict
1619
max_retries: int = 3
17-
retry_delay: int = 10 # seconds
20+
retry_delay: int = 10 # 秒
21+
22+
DEFAULT_API_BASE = "https://api.minimaxi.com/v1"
23+
REQUEST_TIMEOUT = (10, 120) # (连接超时, 读取超时)
24+
MAX_POLL_ATTEMPTS = 60 # 最多轮询 60 次(约 10 分钟)
25+
26+
# V2 (MiniMax-H3) 专用参数
27+
V2_EXTRA_FIELDS = ("resolution", "duration", "ratio", "callback_url")
28+
SUCCESS_STATUSES = frozenset({"succeeded", "Success"})
29+
FAIL_STATUSES = frozenset({"failed", "Fail", "cancelled", "Cancel"})
30+
ERROR_KEYS = ("error_message", "error", "detail", "message", "msg")
1831

1932
def __init__(self, **kwargs):
2033
super().__init__(**kwargs)
2134
self.api_key = kwargs.get("api_key")
22-
self.api_base = kwargs.get("api_base", "https://api.minimaxi.com/v1")
35+
self.api_base = kwargs.get("api_base", self.DEFAULT_API_BASE)
2336
self.model_name = kwargs.get("model_name")
24-
self.params = kwargs.get("params", {})
37+
self.params = kwargs.get("params", {}) or {}
2538
self.max_retries = kwargs.get("max_retries", 3)
26-
self.retry_delay = 10
39+
self.retry_delay = kwargs.get("retry_delay", 10)
40+
# 显式参数可覆盖自动探测(params.api_version: 'v1' / 'v2')
41+
self.api_version = self.params.get("api_version", "auto")
42+
self._session = self._build_session()
2743

2844
@staticmethod
2945
def is_cache_model():
@@ -48,127 +64,194 @@ def new_instance(model_type, model_name, model_credential: Dict[str, object], **
4864
def check_auth(self):
4965
return True
5066

51-
def _safe_call(self, method, url, **kwargs):
52-
"""带重试的请求封装"""
53-
headers = {"Authorization": f"Bearer {self.api_key}"}
54-
55-
for attempt in range(self.max_retries):
67+
def _build_session(self) -> requests.Session:
68+
"""创建带鉴权头与连接复用的请求会话。"""
69+
session = requests.Session()
70+
session.headers.update({"Authorization": f"Bearer {self.api_key}"})
71+
return session
72+
73+
# ---------- API 版本探测 / URL 构建 ----------
74+
75+
def _detect_api_version(self) -> str:
76+
"""探测当前使用 V1 还是 V2 (MiniMax-H3)。"""
77+
if self.api_version in ("v1", "v2"):
78+
return self.api_version
79+
# 模型名包含 H3 -> V2
80+
if self.model_name and "H3" in self.model_name.upper():
81+
return "v2"
82+
# api_base 路径包含 /v2 -> V2
83+
base_path = self.api_base.split("://", 1)[-1] if "://" in self.api_base else self.api_base
84+
if "/v2" in base_path:
85+
return "v2"
86+
return "v1"
87+
88+
def _base_url(self) -> str:
89+
"""去掉结尾的 /v1 或 /v2,返回纯净 base,便于拼装两套路径。"""
90+
base = self.api_base.rstrip("/")
91+
if base.endswith("/v1") or base.endswith("/v2"):
92+
base = base[:-3]
93+
return base.rstrip("/")
94+
95+
def _v2(self) -> bool:
96+
return self._detect_api_version() == "v2"
97+
98+
# ---------- 底层请求 / 轮询 ----------
99+
100+
def _request(self, method: str, url: str, **kwargs) -> dict:
101+
"""带固定间隔重试的请求封装,成功返回 JSON 响应。"""
102+
kwargs.setdefault("timeout", self.REQUEST_TIMEOUT)
103+
for attempt in range(1, self.max_retries + 1):
56104
try:
57-
if method.upper() == "POST":
58-
response = requests.post(url, headers=headers, **kwargs)
59-
elif method.upper() == "GET":
60-
response = requests.get(url, headers=headers, **kwargs)
61-
else:
62-
raise ValueError(f"Unsupported HTTP method: {method}")
63-
105+
response = self._session.request(method, url, **kwargs)
64106
response.raise_for_status()
65107
return response.json()
66108
except (
67109
requests.exceptions.ProxyError,
68110
requests.exceptions.ConnectionError,
69111
requests.exceptions.Timeout,
70-
) as e:
71-
maxkb_logger.error(f"⚠️ 网络错误: {e},正在重试 {attempt + 1}/{self.max_retries}...")
72-
time.sleep(self.retry_delay)
73-
except requests.exceptions.HTTPError as e:
74-
maxkb_logger.error(f"HTTP 错误: {e}")
75-
raise RuntimeError(f"HTTP 请求失败: {e.response.text if hasattr(e, 'response') else str(e)}")
76-
112+
) as exc:
113+
if attempt < self.max_retries:
114+
maxkb_logger.warning(f"网络错误: {exc},正在重试 {attempt + 1}/{self.max_retries}...")
115+
time.sleep(self.retry_delay)
116+
else:
117+
raise RuntimeError("多次重试后仍无法连接到 MiniMax API,请检查代理或网络配置") from exc
118+
except requests.exceptions.HTTPError as exc:
119+
detail = exc.response.text if exc.response is not None else str(exc)
120+
raise RuntimeError(f"HTTP 请求失败: {detail}") from exc
77121
raise RuntimeError("多次重试后仍无法连接到 MiniMax API,请检查代理或网络配置")
78122

123+
def _wait_for_result(self, query_url: str, task_id: Optional[str] = None) -> dict:
124+
"""轮询任务状态直至成功/失败,成功时返回原始响应。"""
125+
params = {"task_id": task_id} if task_id else None
126+
for attempt in range(1, self.MAX_POLL_ATTEMPTS + 1):
127+
response_data = self._request("GET", query_url, params=params)
128+
task = response_data.get("task") or response_data
129+
status = task.get("status")
130+
131+
maxkb_logger.info(f"当前任务状态 (尝试 {attempt}/{self.MAX_POLL_ATTEMPTS}): {status}")
132+
133+
if status in self.SUCCESS_STATUSES:
134+
return response_data
135+
if status in self.FAIL_STATUSES:
136+
error_msg = self._extract_error(task, response_data)
137+
raise RuntimeError(f"视频生成失败: {error_msg}")
138+
# queued / running 等状态,继续轮询
139+
time.sleep(self.retry_delay)
140+
141+
raise RuntimeError(f"任务超时:经过 {self.MAX_POLL_ATTEMPTS} 次轮询后仍未完成")
142+
143+
@staticmethod
144+
def _extract_error(task: dict, response_data: dict) -> str:
145+
for container in (task, response_data):
146+
if not isinstance(container, dict):
147+
continue
148+
for key in GenerationVideoModel.ERROR_KEYS:
149+
value = container.get(key)
150+
if value:
151+
return str(value)
152+
return "未知错误"
153+
154+
# ---------- 对外入口 ----------
155+
79156
def generate_video(self, prompt, negative_prompt=None, first_frame_url=None, last_frame_url=None, **kwargs):
80157
"""
81-
生成视频
158+
生成视频。
82159
prompt: 文本描述
83160
negative_prompt: 反向文本描述(MiniMax 暂不支持,保留参数以兼容接口)
84161
first_frame_url: 起始关键帧图片 URL (图生视频或首尾帧模式)
85162
last_frame_url: 结束关键帧图片 URL (首尾帧模式)
86163
87164
返回: 视频下载 URL
88165
"""
89-
base_url = f"{self.api_base}/video_generation"
166+
# 自动兼容 V1 / V2 (MiniMax-H3) 两套参数逻辑
167+
if self._v2():
168+
return self._generate_video_v2(prompt, first_frame_url, last_frame_url, **kwargs)
169+
return self._generate_video_v1(prompt, first_frame_url, last_frame_url, **kwargs)
170+
171+
# ---------- V2 (MiniMax-H3) 流程 ----------
172+
173+
def _build_v2_payload(self, prompt, first_frame_url, last_frame_url) -> dict:
174+
content = [{"type": "text", "text": prompt}]
175+
if first_frame_url:
176+
content.append(
177+
{
178+
"type": "image_url",
179+
"image_url": {"url": first_frame_url},
180+
"role": "first_frame",
181+
}
182+
)
183+
if last_frame_url:
184+
content.append(
185+
{
186+
"type": "image_url",
187+
"image_url": {"url": last_frame_url},
188+
"role": "last_frame",
189+
}
190+
)
191+
192+
payload = {"model": self.model_name, "content": content}
193+
# V2 必需的 resolution / duration,以及可选的 ratio / callback_url 均来自 params
194+
for key in self.V2_EXTRA_FIELDS:
195+
if key in self.params:
196+
payload[key] = self.params[key]
197+
return payload
198+
199+
def _generate_video_v2(self, prompt, first_frame_url=None, last_frame_url=None, **kwargs) -> str:
200+
base_url = f"{self._base_url()}/v2/video_generation"
201+
payload = self._build_v2_payload(prompt, first_frame_url, last_frame_url)
202+
203+
maxkb_logger.info(f"提交视频生成任务(V2/H3),模型: {self.model_name}")
204+
response_data = self._request("POST", base_url, json=payload)
90205

91-
# 构建基础参数
92-
payload = {
93-
"prompt": prompt,
94-
"model": self.model_name,
95-
}
206+
task_id = response_data.get("task_id")
207+
if not task_id:
208+
raise RuntimeError(f"提交任务失败,未获取到 task_id: {response_data}")
209+
210+
query_url = f"{self._base_url()}/v2/query/video_generation/{task_id}"
211+
response_data = self._wait_for_result(query_url)
212+
213+
task = response_data.get("task") or response_data
214+
video_url = (task.get("content") or {}).get("url")
215+
if not video_url:
216+
raise RuntimeError(f"任务成功但未获取到视频 URL: {response_data}")
217+
return video_url
96218

219+
# ---------- V1 流程(兼容老接口) ----------
220+
221+
def _generate_video_v1(self, prompt, first_frame_url=None, last_frame_url=None, **kwargs) -> str:
222+
base_url = f"{self._base_url()}/v1/video_generation"
223+
224+
payload = {"prompt": prompt, "model": self.model_name}
97225
# 根据提供的参数判断生成模式
98226
if first_frame_url and last_frame_url:
99-
# 模式三:首尾帧生成视频
100-
payload["first_frame_image"] = first_frame_url
101-
payload["last_frame_image"] = last_frame_url
102-
maxkb_logger.info("使用首尾帧模式生成视频")
227+
payload.update(first_frame_image=first_frame_url, last_frame_image=last_frame_url)
103228
elif first_frame_url:
104-
# 模式二:图生视频
105229
payload["first_frame_image"] = first_frame_url
106-
maxkb_logger.info("使用图生视频模式")
107-
else:
108-
# 模式一:文生视频
109-
maxkb_logger.info("使用文生视频模式")
110230

111-
# 合并额外参数(duration, resolution 等)
112-
payload.update(self.params)
231+
# 合并额外参数(duration, resolution 等),跳过版本探测专用字段
232+
payload.update({k: v for k, v in self.params.items() if k != "api_version"})
113233

114-
# --- 步骤 1: 提交任务 ---
115-
maxkb_logger.info(f"提交视频生成任务,模型: {self.model_name}")
116-
response_data = self._safe_call("POST", base_url, json=payload)
234+
maxkb_logger.info(f"提交视频生成任务(V1),模型: {self.model_name}")
235+
response_data = self._request("POST", base_url, json=payload)
117236

118237
task_id = response_data.get("task_id")
119238
if not task_id:
120239
raise RuntimeError(f"提交任务失败,未获取到 task_id: {response_data}")
121240

122-
maxkb_logger.info(f"任务已提交,task_id: {task_id}")
123-
124-
# --- 步骤 2: 轮询查询任务状态 ---
125-
query_url = f"{self.api_base}/query/video_generation"
126-
file_id = self._poll_task_status(query_url, task_id)
241+
query_url = f"{self._base_url()}/v1/query/video_generation"
242+
response_data = self._wait_for_result(query_url, task_id=task_id)
127243

128-
# --- 步骤 3: 获取视频下载链接 ---
129-
video_url = self._get_video_download_url(file_id)
244+
file_id = response_data.get("file_id")
245+
if not file_id:
246+
raise RuntimeError(f"任务成功但未获取到 file_id: {response_data}")
247+
return self._get_video_download_url_v1(file_id)
130248

131-
maxkb_logger.info(f"视频生成完成!视频 URL: {video_url}")
132-
return video_url
133-
134-
def _poll_task_status(self, query_url: str, task_id: str) -> str:
135-
"""轮询任务状态,直至成功或失败"""
136-
params = {"task_id": task_id}
137-
max_attempts = 60 # 最多轮询 60 次(约 10 分钟)
138-
139-
for attempt in range(max_attempts):
140-
response_data = self._safe_call("GET", query_url, params=params)
141-
status = response_data.get("status")
142-
143-
maxkb_logger.info(f"当前任务状态 (尝试 {attempt + 1}/{max_attempts}): {status}")
144-
145-
if status == "Success":
146-
file_id = response_data.get("file_id")
147-
if not file_id:
148-
raise RuntimeError(f"任务成功但未获取到 file_id: {response_data}")
149-
maxkb_logger.info(f"任务处理成功,file_id: {file_id}")
150-
return file_id
151-
elif status == "Fail":
152-
error_msg = response_data.get("error_message", "未知错误")
153-
maxkb_logger.error(f"视频生成失败: {error_msg}")
154-
raise RuntimeError(f"视频生成失败: {error_msg}")
155-
else:
156-
# 任务仍在处理中,等待后继续轮询
157-
time.sleep(self.retry_delay)
158-
159-
raise RuntimeError(f"任务超时:经过 {max_attempts} 次轮询后仍未完成")
160-
161-
def _get_video_download_url(self, file_id: str) -> str:
162-
"""根据 file_id 获取视频下载链接"""
163-
retrieve_url = f"{self.api_base}/files/retrieve"
164-
params = {"file_id": file_id}
165-
166-
response_data = self._safe_call("GET", retrieve_url, params=params)
167-
168-
file_info = response_data.get("file", {})
169-
download_url = file_info.get("download_url")
249+
def _get_video_download_url_v1(self, file_id: str) -> str:
250+
"""根据 file_id 获取视频下载链接(V1)。"""
251+
retrieve_url = f"{self._base_url()}/v1/files/retrieve"
252+
response_data = self._request("GET", retrieve_url, params={"file_id": file_id})
170253

254+
download_url = (response_data.get("file") or {}).get("download_url")
171255
if not download_url:
172256
raise RuntimeError(f"获取下载链接失败: {response_data}")
173-
174257
return download_url

0 commit comments

Comments
 (0)