Files
comfyui_o1key/clients/gemini_flash_client.py
T
o1keyandClaude Sonnet 4.5 9ab209b2b7 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]>
2026-04-03 16:18:45 +08:00

304 lines
9.7 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.
"""
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)