diff --git a/apps/models_provider/impl/minimax_model_provider/model/ttv.py b/apps/models_provider/impl/minimax_model_provider/model/ttv.py index 65eaf45a37f..9f904da422d 100644 --- a/apps/models_provider/impl/minimax_model_provider/model/ttv.py +++ b/apps/models_provider/impl/minimax_model_provider/model/ttv.py @@ -1,5 +1,5 @@ import time -from typing import Dict +from typing import Dict, ClassVar import requests @@ -16,22 +16,20 @@ 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(): @@ -39,16 +37,16 @@ def is_cache_model(): @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, ) @@ -60,26 +58,26 @@ 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): """带重试的请求封装""" @@ -87,18 +85,20 @@ def _safe_call(self, method, url, **kwargs): 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: @@ -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, @@ -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: @@ -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") @@ -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: @@ -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}") @@ -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")