Files
comfyui_o1key/utils/http_error.py
T
Jony ba920f2b66 Publish current ComfyUI O1Key code baseline
Replace the prior release tree with the current plugin, frontend, tests, and documentation. Document retired node IDs and the public Gitea update source.
2026-09-24 19:56:48 +08:00

394 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
统一 HTTP 错误处理 & 退避重试模块
使用方式:
1. 对于 aiohttp 请求,用 async_request_with_retry() 包裹 POST/GET 调用
2. 对于已拿到 status code 的场景,调用 raise_for_status() 抛出友好错误
新增生图/视频节点时,请统一使用本模块处理 HTTP 错误。
"""
import asyncio
import json
import random
import re
from typing import Any, Optional
import aiohttp
# ═══════════════════════════════════════════════════════════════════════════════
# 状态码 → 用户友好文案
# ═══════════════════════════════════════════════════════════════════════════════
HTTP_ERROR_MESSAGES = {
402: "账户余额或模型额度不足,请充值或检查令牌额度。",
429: "模型速率超限或额度不足!",
422: "输入内容未通过安全检查,请调整提示词或参考素材。",
502: "网关超时。请重试或将网络切换为美国直连",
503: "模型超载。请稍后重试!",
504: "网关超时。请稍后重试。",
529: "MiniMax 上游模型过载,请稍后重试。",
}
# 【o1key 图片生成】节点的模型无关错误封装。GPT Image 与所有
# Nano Banana 分支都必须使用同一份映射,不能按模型拆分。
O1KEY_IMAGE_ERROR_CONTENT_MESSAGES = {
"content rejected: the image was flagged as unsafe by the content safety system": "内容被拒绝:该图像被内容安全系统标记为不安全。",
"Your request was rejected by the safety system": "您的请求已被安全系统拒绝",
"insufficient balance": "上游额度不足!",
"Image generation returned empty response": "图片生成过程中被内容审查机制拒绝!",
"The provided prompt is considered unsafe and it cannot be used to generate content": "提供的提示被认为是不安全的,不能用于生成内容。",
}
# 错误内容关键词 → 用户友好文案(优先于状态码匹配)
ERROR_CONTENT_MESSAGES = {
**O1KEY_IMAGE_ERROR_CONTENT_MESSAGES,
"safety system": "请求被安全系统拦截:请调整提示词,避免敏感、违规、血腥、色情、仇恨、未成年人或真实人物等高风险内容。",
"unexpected end of JSON input": "通常重试能解决;反复出现就降低分辨率、数量,或在 API密钥设置中切换全局线路。",
"The current model has a high load": "模型过载,请稍后重试!",
"system error": "系统错误,请稍后重试。",
}
def format_o1key_image_error(value: Any) -> str:
"""Format provider errors for every model in the o1key image generator."""
message = str(value or "生成失败")
message_lower = message.lower()
for keyword, friendly_message in O1KEY_IMAGE_ERROR_CONTENT_MESSAGES.items():
if keyword.lower() in message_lower:
return friendly_message
return message
O1KEY_VIDEO_COPYRIGHT_MESSAGES = {
"audio": "请求失败,输出视频中音频触发版权限制!",
"video": "请求失败,输出视频触发版权限制!",
"content": "请求失败,提示词触发版权限制!",
"real": "请求失败,真人内容触发版权限制!",
"unknown": "请求失败,生成内容触发版权限制!",
}
O1KEY_VIDEO_REVIEW_MESSAGES = {
"audio": "请求失败,输出视频中音频触发审查!",
"video": "请求失败,输出视频触发审查!",
"content": "请求失败,提示词触发审查!",
"real": "请求失败,真人内容触发审查!",
"unknown": "请求失败,生成内容触发审查!",
}
O1KEY_VIDEO_REVIEW_MARKERS = (
"sensitive",
"safety",
"moderation",
"policy violation",
"policy_violation",
"policyviolation",
"unsafe",
"censor",
"review",
)
O1KEY_VIDEO_SUBJECT_REVIEW_MARKERS = ("rejected", "blocked")
_O1KEY_VIDEO_FIELD_RE = re.compile(
r"(?:[\"']?(?:field|type|category|source)[\"']?\s*[:=]\s*[\"']?)"
r"(audio|video|content|real)\b",
re.IGNORECASE,
)
def _o1key_video_error_subject(message: str) -> str:
"""Identify which part of a video request triggered an upstream review."""
explicit_field = _O1KEY_VIDEO_FIELD_RE.search(message)
if explicit_field:
return explicit_field.group(1).lower()
message_lower = message.lower()
if re.search(r"\baudio\b", message_lower):
return "audio"
if re.search(r"\breal\b|\breal[-_ ]?person\b|真人", message_lower):
return "real"
if re.search(r"\bprompt\b|input[-_ ]?content|提示词", message_lower):
return "content"
if "outputvideo" in message_lower or "output_video" in message_lower:
return "video"
if re.search(r"output[-_ ]+video|\bvideo\b", message_lower):
return "video"
if re.search(r"\bcontent\b", message_lower):
return "content"
return "unknown"
def format_o1key_video_error(value: Any) -> str:
"""Format review errors for every model in the o1key video generator.
Upstream providers use several envelope shapes, but their error text or
field values consistently identify the reviewed subject as audio, video,
content (the prompt), or real-person content. Copyright takes precedence
over the broader safety-review markers.
"""
message = str(value or "视频生成失败")
message_lower = message.lower()
subject = _o1key_video_error_subject(message)
if "copyright" in message_lower:
return O1KEY_VIDEO_COPYRIGHT_MESSAGES[subject]
has_review_marker = any(
marker in message_lower for marker in O1KEY_VIDEO_REVIEW_MARKERS
)
has_subject_rejection = subject != "unknown" and any(
marker in message_lower for marker in O1KEY_VIDEO_SUBJECT_REVIEW_MARKERS
)
if has_review_marker or has_subject_rejection:
return O1KEY_VIDEO_REVIEW_MESSAGES[subject]
return message
# 可退避重试的状态码
RETRYABLE_STATUS_CODES = {429, 502, 503, 504, 524, 529}
# 退避重试默认参数
DEFAULT_MAX_RETRIES = 3
DEFAULT_BASE_DELAY = 2.0 # 首次重试等待秒数
DEFAULT_MAX_DELAY = 30.0 # 最大等待秒数
DEFAULT_BACKOFF_FACTOR = 2.0 # 指数退避因子
def _extract_message_from_payload(payload: Any) -> str:
if isinstance(payload, str):
text = payload.strip()
if not text:
return ""
if text.startswith("{") or text.startswith("["):
try:
return _extract_message_from_payload(json.loads(text))
except Exception:
return text
return text
if not isinstance(payload, dict):
return ""
error = payload.get("error")
if isinstance(error, dict):
for key in ("message", "msg", "detail", "reason"):
value = error.get(key)
if value:
return _extract_message_from_payload(value)
elif error:
return _extract_message_from_payload(error)
for key in ("message", "msg", "detail", "reason", "error_message"):
value = payload.get(key)
if value:
return _extract_message_from_payload(value)
for key in ("data", "result", "response", "output"):
value = payload.get(key)
nested = _extract_message_from_payload(value)
if nested:
return nested
return ""
def extract_structured_error_message(raw_message: str) -> str:
if not isinstance(raw_message, str):
return ""
text = raw_message.strip()
if not (text.startswith("{") or text.startswith("[")):
return ""
return _extract_message_from_payload(text)
def get_friendly_message(status_code: int, raw_message: str = "") -> str:
"""根据状态码/错误内容返回友好文案,未匹配则返回原始信息"""
if status_code == 524:
return "Gateway timed out while waiting for upstream image generation. Please retry, lower resolution/count, or switch network route."
if raw_message:
structured_message = extract_structured_error_message(raw_message)
message_for_matching = structured_message or raw_message
raw_message_lower = message_for_matching.lower()
for keyword, friendly_msg in ERROR_CONTENT_MESSAGES.items():
if keyword.lower() in raw_message_lower:
return friendly_msg
if structured_message:
return structured_message
if status_code == 500:
return "服务器返回 500:上游生成失败或服务端临时异常。请稍后重试;如果多次出现,请降低分辨率/数量,或调整提示词。"
friendly = HTTP_ERROR_MESSAGES.get(status_code)
if friendly:
return friendly
return raw_message or f"请求失败 ({status_code})"
def raise_for_status(status_code: int, raw_message: str = "", prefix: str = ""):
"""根据状态码抛出带友好文案的 RuntimeError"""
friendly = get_friendly_message(status_code, raw_message)
full_msg = f"{prefix}{friendly}" if prefix else friendly
raise RuntimeError(full_msg)
def is_retryable(status_code: int) -> bool:
return status_code in RETRYABLE_STATUS_CODES
def extract_error_detail(payload: Any) -> dict:
"""Extract error_detail from async task payloads."""
if not isinstance(payload, dict):
return {}
queue = [payload]
seen = set()
while queue:
current = queue.pop(0)
if not isinstance(current, dict):
continue
obj_id = id(current)
if obj_id in seen:
continue
seen.add(obj_id)
for key in ("error_detail", "errorDetail", "error_details", "errorDetails"):
detail = current.get(key)
if isinstance(detail, dict):
return detail
for key in ("data", "result", "response", "output", "task_result", "content"):
nested = current.get(key)
if isinstance(nested, dict):
queue.append(nested)
return {}
def _coerce_int(value: Any) -> Optional[int]:
if isinstance(value, bool) or value is None:
return None
if isinstance(value, int):
return value
if isinstance(value, float) and value.is_integer():
return int(value)
if isinstance(value, str):
text = value.strip()
if not text:
return None
try:
return int(float(text))
except ValueError:
return None
return None
def extract_error_status_code(error_detail: Any) -> Optional[int]:
"""Extract the real upstream HTTP status from error_detail."""
if not isinstance(error_detail, dict):
return None
for key in (
"upstream_status",
"upstreamStatus",
"upstream_status_code",
"status_code",
"statusCode",
"http_status",
"httpStatus",
"status",
):
status_code = _coerce_int(error_detail.get(key))
if status_code is not None:
return status_code
return None
def is_error_detail_retryable(error_detail: Any) -> bool:
if not isinstance(error_detail, dict):
return False
retryable = error_detail.get("retryable")
if retryable is True:
return True
if isinstance(retryable, str):
return retryable.strip().lower() in ("true", "1", "yes")
return False
def extract_retry_after_seconds(error_detail: Any) -> Optional[float]:
if not isinstance(error_detail, dict):
return None
for key in ("retry_after_seconds", "retryAfterSeconds", "retry_after", "retryAfter"):
value = error_detail.get(key)
if isinstance(value, bool) or value is None:
continue
try:
seconds = float(value)
except (TypeError, ValueError):
continue
if seconds >= 0:
return seconds
return None
def _compute_delay(attempt: int, base_delay: float, max_delay: float, backoff_factor: float) -> float:
"""计算第 attempt 次重试的等待时间(含 jitter)"""
delay = base_delay * (backoff_factor ** attempt)
delay = min(delay, max_delay)
jitter = random.uniform(0, delay * 0.3)
return delay + jitter
async def async_request_with_retry(
session: aiohttp.ClientSession,
method: str,
url: str,
*,
max_retries: int = DEFAULT_MAX_RETRIES,
base_delay: float = DEFAULT_BASE_DELAY,
max_delay: float = DEFAULT_MAX_DELAY,
backoff_factor: float = DEFAULT_BACKOFF_FACTOR,
prefix: str = "",
**request_kwargs,
) -> aiohttp.ClientResponse:
"""
带退避重试的 aiohttp 请求。
仅对 RETRYABLE_STATUS_CODES (429/502/503/504/524) 进行重试。
超过最大重试次数后抛出友好 RuntimeError。
成功时返回 response 对象(调用者需在 async with 外自行处理 body)。
用法示例:
resp = await async_request_with_retry(session, "POST", url, json=body, headers=headers)
data = await resp.json()
"""
last_status: Optional[int] = None
last_message = ""
for attempt in range(max_retries + 1):
try:
resp = await session.request(method, url, **request_kwargs)
except (aiohttp.ClientError, asyncio.TimeoutError) as e:
if attempt < max_retries:
delay = _compute_delay(attempt, base_delay, max_delay, backoff_factor)
print(f"{prefix}网络错误,{delay:.1f}s 后重试 ({attempt+1}/{max_retries})...")
await asyncio.sleep(delay)
continue
raise RuntimeError(f"{prefix}网络错误: {e}") from None
if resp.status == 200:
return resp
last_status = resp.status
try:
last_message = await resp.text()
except Exception:
last_message = ""
if is_retryable(resp.status) and attempt < max_retries:
delay = _compute_delay(attempt, base_delay, max_delay, backoff_factor)
friendly = get_friendly_message(resp.status)
print(f"{prefix}{friendly} {delay:.1f}s 后重试 ({attempt+1}/{max_retries})...")
await asyncio.sleep(delay)
continue
break
if last_status and last_status in HTTP_ERROR_MESSAGES:
raise_for_status(last_status, raw_message=last_message, prefix=prefix)
friendly = get_friendly_message(last_status or 0, last_message)
raise RuntimeError(f"{prefix}{friendly}")