Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 54 additions & 50 deletions apps/models_provider/impl/minimax_model_provider/model/ttv.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import time
from typing import Dict
from typing import Dict, ClassVar

import requests

Expand All @@ -16,39 +16,37 @@ class GenerationVideoModel(MaxKBBaseModel, BaseGenerationVideo):
max_retries: int = 3
retry_delay: int = 10 # seconds

# V2 (MiniMax-H3) 专用参数
v2_extra_fields = ("resolution", "duration", "ratio", "callback_url")
# V2 完成 / 失败状态
v2_success_status = ("succeeded", "Success")
v2_fail_status = ("failed", "Fail", "cancelled", "Cancel")
v2_extra_fields: ClassVar[tuple] = ("resolution", "duration", "ratio", "callback_url")
v2_success_status: ClassVar[frozenset] = frozenset({"succeeded", "Success"})
v2_fail_status: ClassVar[frozenset] = frozenset({"failed", "Fail", "cancelled", "Cancel"})

def __init__(self, **kwargs):
super().__init__(**kwargs)
self.api_key = kwargs.get('api_key')
self.api_base = kwargs.get('api_base', 'https://api.minimaxi.com/v1')
self.model_name = kwargs.get('model_name')
self.params = kwargs.get('params', {}) or {}
self.max_retries = kwargs.get('max_retries', 3)
self.retry_delay = kwargs.get('retry_delay', 10)
self.api_key = kwargs.get("api_key")
self.api_base = kwargs.get("api_base", "https://api.minimaxi.com/v1")
self.model_name = kwargs.get("model_name")
self.params = kwargs.get("params", {}) or {}
self.max_retries = kwargs.get("max_retries", 3)
self.retry_delay = kwargs.get("retry_delay", 10)
# 显式参数可覆盖自动探测(params.api_version: 'v1' / 'v2')
self.api_version = self.params.get('api_version', 'auto')
self.api_version = self.params.get("api_version", "auto")

@staticmethod
def is_cache_model():
return False

@staticmethod
def new_instance(model_type, model_name, model_credential: Dict[str, object], **model_kwargs):
optional_params = {'params': {}}
optional_params = {"params": {}}
for key, value in model_kwargs.items():
if key not in ['model_id', 'use_local', 'streaming']:
optional_params['params'][key] = value
if key not in ["model_id", "use_local", "streaming"]:
optional_params["params"][key] = value

api_base = model_credential.get('api_base', 'https://api.minimaxi.com/v1')
api_base = model_credential.get("api_base", "https://api.minimaxi.com/v1")

return GenerationVideoModel(
model_name=model_name,
api_key=model_credential.get('api_key'),
api_key=model_credential.get("api_key"),
api_base=api_base,
**optional_params,
)
Expand All @@ -60,45 +58,47 @@ def check_auth(self):

def _detect_api_version(self) -> str:
"""探测当前使用 V1 还是 V2 (MiniMax-H3)。"""
if self.api_version in ('v1', 'v2'):
if self.api_version in ("v1", "v2"):
return self.api_version
# 模型名包含 H3 -> V2
if self.model_name and 'H3' in self.model_name.upper():
return 'v2'
if self.model_name and "H3" in self.model_name.upper():
return "v2"
# api_base 路径包含 /v2 -> V2
base_path = self.api_base.split('://', 1)[-1] if '://' in self.api_base else self.api_base
if '/v2' in base_path:
return 'v2'
return 'v1'
base_path = self.api_base.split("://", 1)[-1] if "://" in self.api_base else self.api_base
if "/v2" in base_path:
return "v2"
return "v1"

def _base_url(self) -> str:
"""去掉结尾的 /v1 或 /v2,返回纯净 base,便于拼装两套路径。"""
base = self.api_base.rstrip('/')
if base.endswith('/v1') or base.endswith('/v2'):
base = self.api_base.rstrip("/")
if base.endswith("/v1") or base.endswith("/v2"):
base = base[:-3]
return base.rstrip('/')
return base.rstrip("/")

def _v2(self) -> bool:
return self._detect_api_version() == 'v2'
return self._detect_api_version() == "v2"

def _safe_call(self, method, url, **kwargs):
"""带重试的请求封装"""
headers = {"Authorization": f"Bearer {self.api_key}"}

for attempt in range(self.max_retries):
try:
if method.upper() == 'POST':
if method.upper() == "POST":
response = requests.post(url, headers=headers, **kwargs)
elif method.upper() == 'GET':
elif method.upper() == "GET":
response = requests.get(url, headers=headers, **kwargs)
else:
raise ValueError(f"Unsupported HTTP method: {method}")

response.raise_for_status()
return response.json()
except (requests.exceptions.ProxyError,
requests.exceptions.ConnectionError,
requests.exceptions.Timeout) as e:
except (
requests.exceptions.ProxyError,
requests.exceptions.ConnectionError,
requests.exceptions.Timeout,
) as e:
maxkb_logger.error(f"⚠️ 网络错误: {e},正在重试 {attempt + 1}/{self.max_retries}...")
time.sleep(self.retry_delay)
except requests.exceptions.HTTPError as e:
Expand Down Expand Up @@ -127,17 +127,21 @@ def generate_video(self, prompt, negative_prompt=None, first_frame_url=None, las
def _build_v2_payload(self, prompt, first_frame_url, last_frame_url):
content = [{"type": "text", "text": prompt}]
if first_frame_url:
content.append({
"type": "image_url",
"image_url": {"url": first_frame_url},
"role": "first_frame",
})
content.append(
{
"type": "image_url",
"image_url": {"url": first_frame_url},
"role": "first_frame",
}
)
if last_frame_url:
content.append({
"type": "image_url",
"image_url": {"url": last_frame_url},
"role": "last_frame",
})
content.append(
{
"type": "image_url",
"image_url": {"url": last_frame_url},
"role": "last_frame",
}
)

payload = {
"model": self.model_name,
Expand All @@ -154,7 +158,7 @@ def _generate_video_v2(self, prompt, first_frame_url=None, last_frame_url=None,
payload = self._build_v2_payload(prompt, first_frame_url, last_frame_url)

maxkb_logger.info(f"提交视频生成任务(V2/H3),模型: {self.model_name}")
response_data = self._safe_call('POST', base_url, json=payload)
response_data = self._safe_call("POST", base_url, json=payload)

task_id = response_data.get("task_id")
if not task_id:
Expand All @@ -169,7 +173,7 @@ def _poll_task_status_v2(self, task_id: str) -> str:
max_attempts = 60 # 最多轮询 60 次(约 10 分钟)

for attempt in range(max_attempts):
response_data = self._safe_call('GET', query_url)
response_data = self._safe_call("GET", query_url)
task = response_data.get("task") or response_data
status = task.get("status")

Expand Down Expand Up @@ -225,11 +229,11 @@ def _generate_video_v1(self, prompt, first_frame_url=None, last_frame_url=None,
maxkb_logger.info("使用文生视频模式")

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

# --- 步骤 1: 提交任务 ---
maxkb_logger.info(f"提交视频生成任务,模型: {self.model_name}")
response_data = self._safe_call('POST', base_url, json=payload)
response_data = self._safe_call("POST", base_url, json=payload)

task_id = response_data.get("task_id")
if not task_id:
Expand All @@ -250,7 +254,7 @@ def _poll_task_status_v1(self, query_url: str, task_id: str) -> str:
max_attempts = 60 # 最多轮询 60 次(约 10 分钟)

for attempt in range(max_attempts):
response_data = self._safe_call('GET', query_url, params=params)
response_data = self._safe_call("GET", query_url, params=params)
status = response_data.get("status")

maxkb_logger.info(f"当前任务状态 (尝试 {attempt + 1}/{max_attempts}): {status}")
Expand All @@ -276,7 +280,7 @@ def _get_video_download_url_v1(self, file_id: str) -> str:
retrieve_url = f"{self._base_url()}/v1/files/retrieve"
params = {"file_id": file_id}

response_data = self._safe_call('GET', retrieve_url, params=params)
response_data = self._safe_call("GET", retrieve_url, params=params)

file_info = response_data.get("file", {})
download_url = file_info.get("download_url")
Expand Down
Loading