feat: sync latest local version as authoritative codebase

Complete rewrite/sync of comfyui_o1key custom nodes.
Treat this commit as the current canonical version.

Co-Authored-By: Claude Sonnet 4.5 <[email protected]>
This commit is contained in:
o1key
2026-04-03 16:18:45 +08:00
co-authored by Claude Sonnet 4.5
commit 9ab209b2b7
46 changed files with 16757 additions and 0 deletions
+303
View File
@@ -0,0 +1,303 @@
"""
Gemini Flash API 客户端
用于调用 Gemini 3 Flash 模型进行多模态文本生成
"""
import asyncio
from typing import Any, Dict, List, Optional
import aiohttp
from ..utils.config import get_api_key_or_raise, get_api_base_url
from ..models_config import (
get_flash_model_endpoint,
get_enabled_flash_models,
get_flash_model_thinking_level_value,
)
from .base_client import BaseAPIClient
class GeminiFlashClient(BaseAPIClient):
"""
Gemini Flash API 客户端
用于调用 Gemini 3 Flash 模型进行多模态文本生成
特点:
- 支持图片和视频输入
- 支持动态思考等级端点(不思考/低/中/高)
"""
def __init__(self, api_key: Optional[str] = None):
"""
初始化客户端
Args:
api_key: API 密钥,如果为 None 则从配置文件或环境变量读取
"""
if api_key is None:
api_key = get_api_key_or_raise("O1KEY_API_KEY")
super().__init__(
base_url=get_api_base_url(),
api_key=api_key,
max_request_size=100 * 1024 * 1024 # 100MB
)
def get_endpoint(
self,
model: str = "gemini-3-flash-preview",
**kwargs
) -> str:
"""
获取模型的 API 端点
Args:
model: 模型名称
Returns:
API 端点路径
"""
endpoint = get_flash_model_endpoint(model)
if endpoint is None:
# 回退到第一个启用的模型端点
default_models = get_enabled_flash_models()
if default_models:
endpoint = get_flash_model_endpoint(default_models[0])
if endpoint is None:
raise ValueError(f"无法获取模型 '{model}' 的端点")
return endpoint
def get_http_error_message(self, status_code: int, error_message: str) -> Optional[str]:
"""Gemini 请求 429/503 时返回图中约定的多行错误框文案。"""
if status_code == 429:
return (
"莫慌!该模型暂时超出速率限制啦\n"
"解决方案如下(任意一种):\n"
"1.切换当前模型\n"
"2.前往后台,修改令牌分组"
)
if status_code == 503:
return (
"警报!谷歌服务器当前过载!\n"
"解决方案如下:\n"
"1.摸会儿鱼吧,我也没办法,谷歌会尽快恢复,嘿嘿~\n"
"2.切换其他模型\n"
"3.前往后台,修改令牌分组"
)
return None
def build_request_body(
self,
prompt: str = "",
model: str = "gemini-3-flash-preview",
thinking_level: str = "不思考",
image_data: Optional[List[Dict[str, str]]] = None,
video_data: Optional[Dict[str, str]] = None,
document_data: Optional[Dict[str, str]] = None,
**kwargs
) -> Dict[str, Any]:
"""
构建 API 请求体
Args:
prompt: 用户提示词
model: 模型名称
thinking_level: 思考等级(不思考/低/中/高)- 通过动态端点控制,不需要在请求体中传递
image_data: 图片数据列表,每个元素包含 mime_type 和 data
video_data: 视频数据,包含 mime_type 和 data
document_data: 文档数据,包含 mime_type 和 data
Returns:
请求体字典
"""
parts = []
# 添加文本部分
if prompt:
parts.append({"text": prompt})
# 添加图片部分(如果有)
if image_data:
for img in image_data:
parts.append({
"inline_data": {
"mime_type": img["mime_type"],
"data": img["data"]
}
})
# 添加视频部分(如果有)
if video_data:
parts.append({
"inline_data": {
"mime_type": video_data["mime_type"],
"data": video_data["data"]
}
})
# 添加文档部分(如果有)
if document_data:
parts.append({
"inline_data": {
"mime_type": document_data["mime_type"],
"data": document_data["data"]
}
})
# 构建请求体
request_body = {
"contents": [
{
"parts": parts
}
]
}
# 对于支持 thinkingConfig 的固定端点模型(如 gemini-3-pro-preview
# 通过请求体传递思考等级;动态端点模型(如 gemini-3-flash-preview
# 通过不同 URL 端点控制,无需此字段
thinking_level_value = get_flash_model_thinking_level_value(model, thinking_level)
if thinking_level_value is not None:
request_body["generationConfig"] = {
"thinkingConfig": {
"thinkingLevel": thinking_level_value
}
}
return request_body
def parse_response(self, response: Dict[str, Any]) -> str:
"""
解析 API 响应,提取生成的文本
Args:
response: API 响应字典
Returns:
生成的文本内容
Raises:
RuntimeError: 解析失败或 API 拒绝时
"""
# 检查 candidatesTokenCount
usage_metadata = response.get("usageMetadata", {})
candidates_token_count = usage_metadata.get("candidatesTokenCount", -1)
if candidates_token_count == 0:
raise RuntimeError(
"内容审核拒绝 - candidatesTokenCount = 0\n\n"
"原因:提示词或输入内容包含不适当内容\n"
"建议:检查并调整输入内容"
)
# 检查 finishReason
candidates = response.get("candidates", [])
if candidates:
for candidate in candidates:
finish_reason = candidate.get("finishReason", "")
if finish_reason and finish_reason not in ["STOP", "MAX_TOKENS"]:
reason_messages = {
"PROHIBITED_CONTENT": "违禁内容拒绝",
"SAFETY": "安全过滤器拒绝",
"RECITATION": "版权问题"
}
error_title = reason_messages.get(finish_reason, f"生成异常 ({finish_reason})")
raise RuntimeError(f"{error_title}\n建议:调整输入内容后重试")
# 提取文本内容
text_parts = []
for candidate in candidates:
content = candidate.get("content", {})
parts = content.get("parts", [])
for part in parts:
if "text" in part:
text_parts.append(part["text"])
if not text_parts:
raise RuntimeError("API 响应中未找到生成的文本")
# 合并所有文本部分
return "\n".join(text_parts)
async def generate_async(
self,
prompt: str,
model: str = "gemini-3-flash-preview",
thinking_level: str = "不思考",
image_data: Optional[List[Dict[str, str]]] = None,
video_data: Optional[Dict[str, str]] = None,
document_data: Optional[Dict[str, str]] = None,
session: Optional[aiohttp.ClientSession] = None
) -> str:
"""
异步生成文本
Args:
prompt: 用户提示词
model: 模型名称
thinking_level: 思考等级(不思考/低/中/高)
image_data: 图片数据列表
video_data: 视频数据
document_data: 文档数据
session: aiohttp 会话
Returns:
生成的文本内容
"""
endpoint = self.get_endpoint(model=model)
request_body = self.build_request_body(
prompt=prompt,
model=model,
thinking_level=thinking_level,
image_data=image_data,
video_data=video_data,
document_data=document_data
)
response = await self.request_async(
endpoint,
request_body,
session
)
return self.parse_response(response)
def generate_sync(
self,
prompt: str,
model: str = "gemini-3-flash-preview",
thinking_level: str = "不思考",
image_data: Optional[List[Dict[str, str]]] = None,
video_data: Optional[Dict[str, str]] = None,
document_data: Optional[Dict[str, str]] = None
) -> str:
"""
同步生成文本(用于 ComfyUI 节点)
Args:
prompt: 用户提示词
model: 模型名称
thinking_level: 思考等级(不思考/低/中/高)
image_data: 图片数据列表
video_data: 视频数据
document_data: 文档数据
Returns:
生成的文本内容
"""
coro = self.generate_async(
prompt=prompt,
model=model,
thinking_level=thinking_level,
image_data=image_data,
video_data=video_data,
document_data=document_data
)
return self.run_async_in_thread(coro)