Initial commit: Comfyui_o1key v1.10.0

This commit is contained in:
o1key
2026-02-06 15:56:30 +08:00
commit 9ee29e17d0
25 changed files with 5735 additions and 0 deletions
+634
View File
@@ -0,0 +1,634 @@
# Comfyui_o1key 开发指南
## 对话原则
始终使用中文进行对话。
## 项目概述
这是一个 ComfyUI 自定义节点插件,通过 api.o1key.com 调用 AI 模型进行图像生成。
### 技术栈
- Python 3.7+
- ComfyUI 框架
- aiohttp (异步 HTTP)
- Pillow (图像处理)
- PyTorch (张量处理)
---
## 目录结构
```
Comfyui_o1key/
├── __init__.py # 节点注册入口
├── models_config.py # 模型配置中心 ⭐ 管理所有支持的模型
├── version.txt # 版本号文件
├── update.bat # Windows 自动更新脚本
├── update.sh # Linux/Mac 自动更新脚本
├── nodes/ # 节点模块
│ ├── __init__.py
│ ├── nano_banana_pro.py # NanoBananaPro 节点
│ └── batch_nano_banana_pro.py # 批量节点
├── utils/ # 工具模块
│ ├── __init__.py
│ ├── image_utils.py # 图像转换工具
│ ├── config.py # 配置管理
│ └── update_checker.py # 更新检查器
├── clients/ # API 客户端
│ ├── __init__.py
│ ├── base_client.py # 客户端基类
│ └── gemini_client.py # Gemini API 客户端
├── .config # API 配置文件(不提交)
├── .config.example # 配置示例
├── requirements.txt # 依赖包
└── README.md # 用户文档
```
---
## 模型管理系统
### 概述
所有 Nano Banana Pro 支持的模型都在 `models_config.py` 中统一管理。要添加新模型或临时关闭某个模型,只需编辑这个文件即可。
### 模型配置文件 (models_config.py)
#### 配置结构
```python
GEMINI_MODELS = [
{
"id": "gemini-3-pro-image-preview-url",
"description": "URL 模式,根据分辨率自动选择端点 (1K/2K/4K)",
"enabled": True,
"endpoint_type": "dynamic",
"endpoint": None # 动态端点,由代码根据分辨率选择
},
{
"id": "gemini-3-pro-image-preview",
"description": "标准模式,固定端点",
"enabled": True,
"endpoint_type": "standard",
"endpoint": "/v1beta/models/gemini-3-pro-image-preview:generateContent"
},
# 更多模型...
]
```
#### 字段说明
| 字段 | 类型 | 必需 | 说明 |
|------|------|------|------|
| `id` | string | 是 | 模型标识符,用于 API 调用 |
| `description` | string | 是 | 模型描述,说明特点和适用场景 |
| `enabled` | boolean | 是 | 是否启用该模型(false 则在节点中隐藏) |
| `endpoint_type` | string | 是 | 端点类型:"dynamic", "standard", "flatfee" |
| `endpoint` | string | 是 | API 端点路径(动态端点设为 None) |
#### 端点类型说明
- **dynamic**: 根据分辨率动态选择端点(如 gemini-3-pro-image-preview-url
- **standard**: 使用固定端点(如 gemini-3-pro-image-preview
- **flatfee**: 固定费用模式端点(如 gemini-3-pro-image-preview-flatfee
### 常见操作
#### 1. 添加新模型
在 `GEMINI_MODELS` 列表末尾添加新模型:
```python
GEMINI_MODELS = [
# ... 现有模型 ...
{
"id": "gemini-新模型名称",
"description": "新模型的描述和特点",
"enabled": True,
"endpoint_type": "standard", # 根据实际情况选择
"endpoint": "/v1beta/models/gemini-新模型名称:generateContent" # 配置端点
}
]
```
**注意**
- **固定端点模型**:直接在 `endpoint` 字段填写完整的端点路径即可,无需修改代码
- **动态端点模型**:如果模型需要根据分辨率动态选择端点,设置 `endpoint_type: "dynamic"` 和 `endpoint: None`,并在 `gemini_client.py` 的 `get_endpoint()` 方法中添加对应逻辑
#### 2. 临时关闭模型
将模型的 `enabled` 字段设为 `False`
```python
{
"id": "gemini-3-pro-image-preview-url",
"description": "URL 模式",
"enabled": False, # 临时关闭
"endpoint_type": "dynamic"
}
```
关闭后,该模型将不会出现在 ComfyUI 节点的下拉列表中。
#### 3. 重新启用模型
将 `enabled` 改回 `True`
```python
{
"id": "gemini-3-pro-image-preview-url",
"enabled": True, # 重新启用
# ...
}
```
#### 4. 修改模型描述
直接编辑 `description` 字段:
```python
{
"id": "gemini-3-pro-image-preview",
"description": "标准模式,固定端点,适用于常规图像生成", # 更新描述
# ...
}
```
### 工具函数
`models_config.py` 提供了一些工具函数,可在代码中使用:
```python
from ..models_config import (
get_enabled_models, # 获取启用的模型列表
get_all_models, # 获取所有模型(包括禁用的)
get_model_config, # 获取指定模型的完整配置
is_model_enabled, # 检查模型是否启用
get_model_description, # 获取模型描述
get_endpoint_type, # 获取端点类型
get_model_endpoint # 获取模型端点
)
# 示例:获取启用的模型
enabled = get_enabled_models()
# ['gemini-3-pro-image-preview-url', 'gemini-3-pro-image-preview', ...]
# 示例:获取模型配置
config = get_model_config("gemini-3-pro-image-preview-url")
# {'id': '...', 'description': '...', 'enabled': True, 'endpoint_type': 'dynamic', 'endpoint': None}
# 示例:获取模型端点
endpoint = get_model_endpoint("gemini-3-pro-image-preview")
# '/v1beta/models/gemini-3-pro-image-preview:generateContent'
```
### 节点集成
所有使用模型列表的节点都会自动从 `models_config.py` 加载:
```python
from ..models_config import get_enabled_models
class NanoBananaPro:
@classmethod
def INPUT_TYPES(cls):
# 自动从配置加载启用的模型
enabled_models = get_enabled_models()
return {
"required": {
"模型": (enabled_models, {
"default": enabled_models[0]
}),
# ...
}
}
```
### 配置验证
`models_config.py` 在加载时会自动验证配置:
- 检查每个模型是否有必需字段(id, description, enabled, endpoint_type, endpoint
- 检查 `endpoint_type` 是否合法(dynamic, standard, flatfee
- 检查非动态端点模型必须配置有效的 `endpoint`
- 检查端点格式是否正确(应以 `/v1beta/models/` 开头)
- 确保至少有一个模型是启用的
如果配置不合法,会在终端打印警告信息。
### 最佳实践
1. **添加新模型前**
- 确认模型使用 Gemini 原生接口格式
- 确认端点规则(dynamic/standard/flatfee
- 编写清晰的描述说明
2. **临时测试**
- 关闭其他模型,只启用测试模型
- 验证功能后再重新启用其他模型
3. **版本控制**
- `models_config.py` 应纳入版本控制
- 重大模型变更应记录在 `CHANGELOG.md` 中
4. **文档更新**
- 添加新模型后,更新 `README.md` 中的模型列表
- 如有特殊使用说明,添加到文档中
---
## 开发新节点流程
### 1. 创建节点文件
在 `nodes/` 目录下创建新的 Python 文件:
```python
# nodes/my_new_node.py
from typing import Optional, Tuple
import torch
from ..utils.image_utils import tensor_to_pil, pil_to_tensor
from ..clients.gemini_client import GeminiAPIClient
class MyNewNode:
"""节点描述"""
def __init__(self):
self.client = None
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"prompt": ("STRING", {"default": "", "multiline": True}),
# 更多参数...
},
"optional": {
"images": ("IMAGE",)
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "execute"
CATEGORY = "image/generation"
def execute(self, prompt: str, images: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor]:
# 实现逻辑
pass
```
### 2. 注册节点
在 `nodes/__init__.py` 中添加导出:
```python
from .my_new_node import MyNewNode
__all__ = ['NanoBananaPro', 'MyNewNode']
```
在根 `__init__.py` 中注册:
```python
from .nodes import NanoBananaPro, MyNewNode
NODE_CLASS_MAPPINGS = {
"NanoBananaPro": NanoBananaPro,
"MyNewNode": MyNewNode
}
NODE_DISPLAY_NAME_MAPPINGS = {
"NanoBananaPro": "Nano Banana Pro",
"MyNewNode": "My New Node"
}
```
### 3. 更新 CHANGELOG.md
记录新增功能。
---
## ComfyUI 节点规范
### INPUT_TYPES 参数类型
| 类型 | 格式 | 示例 |
|------|------|------|
| 字符串 | `("STRING", {...})` | `("STRING", {"default": "", "multiline": True})` |
| 整数 | `("INT", {...})` | `("INT", {"default": 1, "min": 1, "max": 100})` |
| 浮点数 | `("FLOAT", {...})` | `("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.1})` |
| 下拉选项 | `([...], {...})` | `(["option1", "option2"], {"default": "option1"})` |
| 图像 | `("IMAGE",)` | 放在 optional 中 |
### 返回值规范
```python
RETURN_TYPES = ("IMAGE", "MASK", "STRING") # 类型元组
RETURN_NAMES = ("images", "mask", "text") # 名称元组
```
### 必须的类属性
```python
FUNCTION = "execute" # 执行函数名
CATEGORY = "image/generation" # 节点分类路径
```
---
## 工具模块使用
### 图像转换 (utils/image_utils.py)
```python
from ..utils.image_utils import tensor_to_pil, pil_to_tensor
# ComfyUI Tensor → PIL Image 列表
pil_images = tensor_to_pil(tensor) # tensor: [B, H, W, C], range [0, 1]
# PIL Image 列表 → ComfyUI Tensor
tensor = pil_to_tensor(pil_images) # 返回 [B, H, W, C], range [0, 1]
# PIL → Base64
from ..utils.image_utils import encode_image_to_base64
b64_str = encode_image_to_base64(pil_image)
# Base64 → PIL
from ..utils.image_utils import decode_base64_to_pil
pil_image = decode_base64_to_pil(b64_str)
```
### 配置管理 (utils/config.py)
```python
from ..utils.config import get_api_key, get_api_key_or_raise, load_config
# 获取 API 密钥(返回 None 如果未找到)
api_key = get_api_key("O1KEY_API_KEY")
# 获取 API 密钥(抛出异常如果未找到)
api_key = get_api_key_or_raise("O1KEY_API_KEY")
# 加载完整配置
config = load_config()
```
---
## API 客户端使用
### 使用 GeminiAPIClient
```python
from ..clients.gemini_client import GeminiAPIClient
# 初始化(自动读取配置)
client = GeminiAPIClient()
# 同步生成(用于 ComfyUI 节点)
images = client.generate_sync(
prompt="描述文字",
model="gemini-3-pro-image-preview-url",
resolution="2K",
aspect_ratio="1:1",
batch_size=1,
images=None, # 可选:输入图像列表
progress_callback=None
)
```
### 创建新的 API 客户端
继承 `BaseAPIClient` 并实现抽象方法:
```python
from ..clients.base_client import BaseAPIClient
class MyAPIClient(BaseAPIClient):
def __init__(self):
super().__init__(
base_url="https://api.example.com",
api_key=get_api_key_or_raise("MY_API_KEY"),
max_request_size=20 * 1024 * 1024
)
def get_endpoint(self, **kwargs) -> str:
return "/v1/generate"
def build_request_body(self, **kwargs) -> dict:
return {"prompt": kwargs.get("prompt", "")}
def parse_response(self, response: dict) -> Any:
return response.get("result")
```
---
## API 端点说明
### Gemini 模型端点
**gemini-3-pro-image-preview-url** (根据分辨率动态选择):
- 1K: `/v1beta/models/gemini-3-pro-image-preview-url:generateContent`
- 2K: `/v1beta/models/gemini-3-pro-image-preview-2k-url:generateContent`
- 4K: `/v1beta/models/gemini-3-pro-image-preview-4k-url:generateContent`
**gemini-3-pro-image-preview** (固定端点):
- `/v1beta/models/gemini-3-pro-image-preview:generateContent`
**gemini-3-pro-image-preview-flatfee** (固定端点):
- `/v1beta/models/gemini-3-pro-image-preview-flatfee:generateContent`
### 请求格式
```json
{
"contents": [{
"role": "user",
"parts": [
{"text": "提示词"},
{"inline_data": {"mime_type": "image/png", "data": "base64..."}}
]
}],
"generationConfig": {
"responseModalities": ["TEXT", "IMAGE"],
"imageConfig": {
"aspectRatio": "1:1",
"imageSize": "2K"
}
}
}
```
---
## 代码规范
### 命名约定
- 类名:PascalCase(如 `NanoBananaPro`
- 函数/方法:snake_case(如 `tensor_to_pil`
- 常量:UPPER_CASE(如 `API_BASE_URL`
- 私有方法:前缀下划线(如 `_load_config`
### 类型注解
所有公开函数必须有类型注解:
```python
def function_name(param1: str, param2: Optional[int] = None) -> List[Image.Image]:
pass
```
### 文档字符串
使用 Google 风格的 docstring
```python
def function_name(param1: str, param2: int) -> bool:
"""
函数简短描述
Args:
param1: 参数1说明
param2: 参数2说明
Returns:
返回值说明
Raises:
ValueError: 异常情况说明
Example:
>>> result = function_name("test", 42)
>>> print(result)
True
"""
pass
```
### 错误处理
```python
try:
# 业务逻辑
pass
except ValueError as e:
# 用户输入错误
print(f"节点名: 输入错误 - {str(e)}")
raise
except RuntimeError as e:
# API 或网络错误
print(f"节点名: API 错误 - {str(e)}")
raise
except Exception as e:
# 未知错误
print(f"节点名: 未知错误 - {str(e)}")
raise
```
---
## 限制与约束
| 限制项 | 值 | 说明 |
|--------|-----|------|
| 请求体大小 | 20MB | 超过会报错 |
| 输入图像数量 | 14张 | 图生图模式限制 |
| 批次大小 | 1-1000 | 并发生成数量 |
| 支持的分辨率 | 1K/2K/4K | API 限制 |
---
## 测试检查清单
新节点开发完成后,验证以下场景:
- [ ] 文生图基础功能
- [ ] 图生图功能(如支持)
- [ ] 不同分辨率(1K/2K/4K
- [ ] 不同宽高比
- [ ] 批量生成
- [ ] 错误处理(无 API 密钥、网络错误等)
- [ ] 边界条件(最大图像数、最大批次)
---
## 更新日志
修改代码后,更新 `CHANGELOG.md` 记录变更。
格式:
```markdown
## [版本号] - 日期
### Added
- 新增功能
### Changed
- 变更内容
### Fixed
- 修复问题
```
---
## 版本发布流程
### 1. 准备发布
发布新版本前确认以下事项:
- [ ] 所有功能测试通过
- [ ] 更新 `CHANGELOG.md`(记录本次变更)
- [ ] 更新 `version.txt`(更新版本号)
- [ ] 更新 `README.md`(如有新功能需要说明)
### 2. 版本号规范
遵循语义化版本 (Semantic Versioning)
- **主版本号** (Major): 重大架构变更、不兼容的 API 修改
- **次版本号** (Minor): 新增功能、向后兼容
- **修订号** (Patch): Bug 修复、小改进
示例:`v1.10.2` → Major.Minor.Patch
### 3. 发布步骤
```bash
# 1. 更新版本号
echo "v1.11.0" > version.txt
# 2. 提交变更
git add .
git commit -m "Release v1.11.0: 添加新功能描述"
# 3. 创建标签
git tag v1.11.0
# 4. 推送到远程
git push origin main --tags
```
### 4. 用户更新
用户运行更新脚本即可获取最新版本:
- **Windows**: 双击 `update.bat`
- **Linux/Mac**: 运行 `./update.sh`
更新脚本会自动:
- 检查远程更新
- 备份配置文件
- 拉取最新代码
- 更新依赖包
- 显示更新日志
---
+26
View File
@@ -0,0 +1,26 @@
# 隐私文件(已弃用配置文件,改用环境变量)
# .config
# Python 缓存
__pycache__/
*.py[cod]
*$py.class
*.so
.Python
# 环境
.env
.venv
env/
venv/
# IDE
.vscode/
.idea/
*.swp
*.swo
*~
# OS
.DS_Store
Thumbs.db
+610
View File
@@ -0,0 +1,610 @@
# Changelog
本项目的所有重要变更都将记录在此文件中。
格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/)。
---
## [1.10.0] - 2026-02-06
### Added ⭐
- **自动更新系统** - 让用户轻松更新插件到最新版本
- 新增 `update.bat` - Windows 自动更新脚本
- 新增 `update.sh` - Linux/Mac 自动更新脚本
- 新增 `version.txt` - 版本号管理文件
- 新增 `utils/update_checker.py` - 启动时自动检查更新
- 新增更新检查功能:每次启动 ComfyUI 时自动检测是否有新版本
- **更新脚本功能**
- ✅ 自动检查远程更新
- ✅ 自动备份和恢复 `.config` 配置文件
- ✅ 自动拉取最新代码
- ✅ 自动更新 Python 依赖包
- ✅ 显示版本变更信息
- ✅ 显示最近更新日志(前 20 行)
- ✅ 友好的彩色终端输出(Linux/Mac)
- ✅ 完善的错误处理和提示
### Changed
- **插件启动流程** (`__init__.py`)
- 集成更新检查模块
- 启动时自动检查是否有新版本
- 如有更新,终端显示友好的更新提示
- 静默失败机制,不影响插件正常加载
- **文档更新** (`README.md`)
- 新增"🔄 更新插件"章节
- 提供两种更新方法:自动更新(推荐)和手动更新
- 详细的跨平台更新说明
- 更新提示和注意事项
### Benefits
- 🎯 **用户友好** - 一键更新,无需手动操作 Git
- 🔒 **配置安全** - 自动备份恢复配置,不会丢失设置
-**依赖同步** - 自动更新 Python 包,确保兼容性
- 📋 **信息透明** - 显示版本变更和更新日志
- 🌍 **跨平台** - 支持 Windows/Linux/Mac
- 🛡️ **稳定可靠** - 完善的错误处理,不影响插件运行
---
## [1.9.1] - 2026-02-05
### Security 🔐
- **错误信息伪装 - 极致隐蔽性**
- 未配置情况下的错误信息伪装成商业授权提示
- 前端弹框显示:`未授权!`
- 控制台日志显示:`请联系作者授权后方可使用!`
- 完全隐藏真实的配置方式(环境变量)
- 让外部用户以为这是需要购买授权的商业插件
### Changed
- **配置模块** (`utils/config.py`)
- 将未找到 API 密钥的错误消息改为简洁的"未授权!"
- 移除所有暴露配置方式的信息
- **节点异常处理** (所有节点)
- `NanoBananaPro` - 添加授权错误检测,特殊处理授权失败日志
- `BatchNanoBananaPro` - 添加授权错误检测,特殊处理授权失败日志
- `GoogleGemini` - 添加授权错误检测,特殊处理授权失败日志
- 其他 ValueError 类型仍正常显示详细错误信息
### Benefits
-**极致隐蔽** - 完全看不出需要配置环境变量
-**伪装商业** - 外部用户以为需要付费授权
-**信息安全** - 不泄露任何配置相关信息
-**功能完整** - 其他错误类型仍正常提示
-**向后兼容** - 已配置用户完全不受影响
---
## [1.9.0] - 2026-02-05
### Security 🔐
- **配置方式优化 - 提升隐蔽性和安全性**
- 移除所有配置相关文件(`setup_config.ps1`, `.config`, `.config.example`
- 改为使用系统环境变量存储 API 密钥
- 避免在项目中留下配置痕迹,提高隐蔽性
- 配置方式对外部用户完全不可见
### Changed
- **配置管理模块重构** (`utils/config.py`)
- 调整读取优先级:环境变量优先 > .config 文件(向后兼容)
- 简化错误提示:仅提示设置环境变量,不再提及配置文件
- 更新模块说明:从"处理 .config 文件"改为"处理环境变量"
- **文档更新**
- `README.md` - 简化配置说明,仅保留环境变量设置方法
- `批量节点使用指南.md` - 更新常见问题中的配置说明
- `.gitignore` - 注释配置文件规则(已弃用)
### Benefits
-**高隐蔽性** - 项目中无任何配置相关文件
-**高安全性** - 敏感信息存储在系统级别,不在项目目录
-**简化维护** - 一行命令创建/更新/删除配置
-**多项目共享** - 环境变量可被其他项目复用
-**向后兼容** - 仍支持从 .config 文件读取(如果存在)
---
## [1.8.0] - 2026-02-04
### Added
- **集成 ComfyUI 原生进度条** 🎉
- `NanoBananaPro` 节点现在支持 UI 绿色进度条显示
- `BatchNanoBananaPro` 节点现在支持 UI 绿色进度条显示
- 使用 `comfy.utils.ProgressBar` 实现实时进度更新
- 进度条在节点上方显示,从 0% 平滑更新到 100%
- 兼容性检查:如果 ProgressBar 不可用,自动降级到终端进度显示
### Improved
- **用户体验提升**
- 生图过程中可视化进度反馈更直观
- 节点运行时自动显示绿色边框(ComfyUI 原生)
- 保留详细的终端进度日志,方便调试
- 批量提示词模式下进度条总数自动调整
### Technical
-`nodes/nano_banana_pro.py` 中集成 ProgressBar
- 创建进度条实例:`ProgressBar(生图数量)`
-`progress_callback` 中调用 `pbar.update(1)`
- 批量提示词模式重新创建进度条以匹配实际总数
-`nodes/batch_nano_banana_pro.py` 中集成 ProgressBar
-`pbar` 参数传递给 `_process_batch_async` 方法
- 每完成一个任务立即更新进度条
- 支持大批量任务的实时进度显示
- 添加 `PROGRESS_BAR_AVAILABLE` 标志进行兼容性检测
### Impact
- ✅ 所有图像生成节点现在都有 UI 进度条
- ✅ 后续新增的视频生成节点可直接复用此实现
- ✅ 不影响现有功能,完全向后兼容
---
## [1.7.0] - 2026-02-03
### Added
- **新增 Google Gemini 节点** (`GoogleGemini`)
- 用于调用 Gemini 3 Flash 模型进行多模态文本生成
- **输入支持**
- 提示词(必填):用户提示词,支持多行
- 系统指令(可选):系统级指令,引导模型行为
- 思考深度:不思考(默认)/ 高
- 图片(可选):支持 ComfyUI IMAGE 类型输入
- 视频(可选):支持 ComfyUI VIDEO 类型输入
- **输出**:文本内容(STRING 类型)
- **支持的视频格式**mp4, mpeg, mov, avi, flv, webm, wmv, 3gpp
- **端点映射**
- 不思考 → `/v1beta/models/gemini-3-flash-preview-nothinking:generateContent`
- 高 → `/v1beta/models/gemini-3-flash-preview-high:generateContent`
- **新增 Gemini Flash API 客户端** (`clients/gemini_flash_client.py`)
- 继承 `BaseAPIClient` 基类
- 支持系统指令配置
- 支持图片和视频的 base64 编码发送
- 智能超时设置(视频请求 5 分钟,其他 3 分钟)
### Technical
- 新增 `GeminiFlashClient` 类处理 Gemini Flash 模型 API 调用
- 新增 `GoogleGemini` 节点类实现多模态文本生成
- 支持自动检测视频 MIME 类型
- 视频文件大小限制 20MB
---
## [1.6.2] - 2026-02-02
### Added
- **模型端点配置化** (`models_config.py`)
- 在模型配置中新增 `endpoint` 字段,集中管理每个模型的 API 端点
- 新增 `get_model_endpoint()` 工具函数,用于获取模型端点
- 添加新模型时只需在配置文件中填写端点,无需修改代码
### Changed
- **简化端点获取逻辑** (`clients/gemini_client.py`)
- 重构 `get_endpoint()` 方法,从配置文件读取端点而非硬编码
- 保留动态端点模型(gemini-3-pro-image-preview-url)的特殊处理逻辑
- 其他模型自动从 `models_config.py` 读取端点配置
### Improved
- **增强配置验证** (`models_config.py`)
- 验证非动态端点模型必须配置有效的 `endpoint`
- 验证端点格式是否正确(应以 `/v1beta/models/` 开头)
- 更新必需字段列表,包含 `endpoint` 字段
- **更新开发文档** (`.cursorrules`)
- 更新模型配置示例,说明 `endpoint` 字段的使用方法
- 更新添加新模型的指南,强调配置端点的方式
- 更新工具函数列表,添加 `get_model_endpoint()` 说明
- 更新配置验证规则说明
### Benefits
- 降低维护成本:添加新模型只需修改配置文件
- 提高可读性:端点集中管理,一目了然
- 减少错误:配置验证确保端点格式正确
- 保持灵活性:动态端点模型仍可使用代码逻辑
---
## [1.6.1] - 2026-02-02
### Fixed
- **修复 gemini-3-pro-image-preview-flatfee 模型 504 错误**
- 暂时禁用 `gemini-3-pro-image-preview-flatfee` 模型(端点返回 504 Gateway Timeout
- 更新模型描述标注"暂时不可用-504错误"
- 建议用户使用其他可用模型(如 gemini-3-pro-image-preview 或 gemini-3-pro-image-preview-url
### Improved
- **增强 HTTP 错误处理** (`clients/base_client.py`)
- 新增针对 504 Gateway Timeout 的友好错误提示
- 说明原因:服务器响应超时或端点暂时不可用
- 提供解决建议:尝试其他模型、稍后重试、降低分辨率等
- 新增针对 503 Service Unavailable 的错误提示
- 说明原因:模型服务过载或维护中
- 提供解决建议:稍后重试或尝试其他模型
- 新增针对 429 Too Many Requests 的错误提示
- 说明原因:API 配额用尽或请求过于频繁
- 提供解决建议:等待后重试或检查配额
- 新增针对 404 Not Found 的错误提示
- 说明原因:端点路径错误或模型不存在
- 提供解决建议:检查模型名称或使用其他模型
- 统一错误信息格式:错误类型 + 原因 + 建议
- 同时优化 POST 和 GET 请求的错误处理
### Technical
-`request_async()``request_get_async()` 中添加状态码判断逻辑
- 提供更详细的错误诊断信息,帮助用户快速定位和解决问题
---
## [1.6.0] - 2026-02-02
### Added
- **模型管理系统** (`models_config.py`)
- 创建集中式模型配置文件,所有 Nano Banana Pro 支持的模型统一管理
- 支持快速添加新模型、临时关闭或启用模型
- 模型配置包含:模型ID、描述、启用状态、端点类型
- 提供工具函数:
- `get_enabled_models()` - 获取启用的模型列表
- `get_all_models()` - 获取所有模型(包括禁用的)
- `get_model_config()` - 获取指定模型的完整配置
- `is_model_enabled()` - 检查模型是否启用
- `get_model_description()` - 获取模型描述
- `get_endpoint_type()` - 获取端点类型
- 自动配置验证,确保配置完整性和合法性
### Changed
- **NanoBananaPro 节点重构**
- 移除硬编码的 `MODELS` 列表
- 改为从 `models_config.py` 动态加载模型列表
- 节点在 ComfyUI 中显示的模型列表自动同步配置文件
- **BatchNanoBananaPro 节点重构**
- 移除硬编码的 `MODELS` 列表
- 改为从 `models_config.py` 动态加载模型列表
- 保持与 NanoBananaPro 节点的模型列表一致性
### Improved
- **开发指南更新** (`.cursorrules`)
- 新增"模型管理系统"章节
- 详细说明模型配置文件结构和字段含义
- 提供添加新模型、关闭/启用模型、修改描述等常见操作指南
- 新增工具函数使用示例和节点集成说明
- 补充配置验证机制和最佳实践建议
- **目录结构更新**
- 在开发指南中添加 `models_config.py` 文件说明
- 标记为模型配置中心 ⭐
### Benefits
-**集中管理** - 所有模型定义在一个文件,易于维护
-**易于扩展** - 添加新模型只需在配置文件中添加一个字典
-**快速开关** - 修改 `enabled` 字段即可临时关闭或启用模型
-**文档化** - 每个模型都有 `description` 说明特点和适用场景
-**类型安全** - Python 文件支持代码提示和类型检查
-**自动同步** - 所有节点自动使用最新的模型配置
---
## [1.5.7] - 2026-02-02
### Performance 🚀
- **极致性能优化:彻底解决异步阻塞问题**
-`parse_response()` 改为 `parse_response_async()`,实现完全异步的图片下载
- 使用 `aiohttp` 替代同步的 `requests.get()` 下载图片
- **问题**:之前在异步事件循环中使用同步 HTTP 请求会阻塞整个事件循环
- **影响**:虽然 API 请求是并发的,但图片下载变成了串行操作
- **效果**:现在图片下载也是完全并发的,真正实现端到端的异步性能
- **性能提升幅度**
- 使用 `gemini-3-pro-image-preview-url` 模型时提升最明显
- 批量生成 4 张图时,从"串行下载 4 张"变为"并发下载 4 张"
- 预计性能提升 2-4 倍(取决于网络延迟和图片大小)
### Changed
- `GeminiAPIClient.parse_response()``parse_response_async()`
- 新增 `session` 参数,用于复用 aiohttp 会话
- 支持并发下载多个图片 URL
- 自动管理 session 生命周期
- `GeminiAPIClient.generate_single_async()` 现在调用异步解析方法
- 移除 `requests` 依赖,统一使用 `aiohttp`
### Impact
- 所有节点自动受益:
-`NanoBananaPro` - 批量生成速度显著提升
-`BatchNanoBananaPro` - 大批量任务性能大幅改善
- ✅ 所有使用 `gemini-3-pro-image-preview-url` 模型的场景
---
## [1.5.6] - 2026-02-02
### Improved
- **优化批量生成实时进度显示** (`Nano Banana Pro`)
- 使用 `asyncio.as_completed` 替代 `asyncio.gather`,实现真正的实时进度
- 每完成一个请求立即显示进度,而非等待所有请求完成后批量显示
- 进度信息包含成功/失败状态:
- 成功:`✓ [1/4] 第 1 张生成成功`
- 失败:`✗ [2/4] 生成失败 - 错误原因`
- 最终统计显示成功和失败数量
### Fixed
- **修复批量生成失败信息丢失问题**
- 之前:失败的请求被静默忽略,用户不知道哪些请求失败
- 现在:每个失败的请求都会显示错误原因,方便排查问题
### Technical
- 更新 `generate_batch_async()``generate_multi_prompts_async()` 方法
- 进度回调签名变更:`(current, total)``(current, total, success, error_msg)`
- 错误信息自动截取第一行,避免过长输出
---
## [1.5.5] - 2026-02-02
### Added
- **增强 API 错误处理机制**
- 新增 `candidatesTokenCount = 0` 检测(最高优先级)
- 自动检测内容审核拒绝情况
- 提供明确的拒绝原因和改进建议
- 新增 `finishReason` 异常检测(次优先级)
- 支持检测 `PROHIBITED_CONTENT`(违禁内容)
- 支持检测 `SAFETY`(安全过滤器)
- 支持检测 `RECITATION`(版权问题)
- 支持检测 `MAX_TOKENS`Token 超限)
- 针对每种错误类型提供具体的解决建议
- 新增 API 文本响应拒绝检测
- 当 API 返回文本而非图像时,自动提取拒绝说明
- 直接展示 API 的拒绝理由给用户
### Improved
- **优化错误提示格式** (`clients/gemini_client.py`)
- 统一错误信息格式:错误类型 + 原因 + 建议
- 所有错误以 `RuntimeError` 抛出,便于节点层面捕获
- 提供清晰的多行格式化错误信息
- 包含针对性的操作建议,帮助用户快速解决问题
### Technical
-`GeminiAPIClient.parse_response()` 方法中实现三层错误检测
- 错误检测按优先级顺序执行,确保最重要的问题优先报告
- 保持向后兼容,不影响正常的图像生成流程
---
## [1.5.4] - 2026-02-01
### Added
- **新增余额查询功能**
- 在每次图像生成请求完成后自动查询并显示用户余额
- 支持查询 API 名称和当前可用余额
- 余额格式:`当前余额:$XX.XX | APIxxx`
- 查询失败时显示警告信息,不影响主流程
- 适用于 `NanoBananaPro``BatchNanoBananaPro` 节点
### Changed
- **扩展 API 客户端功能** (`clients/base_client.py`)
- 新增 `request_get_async()` 方法支持 GET 请求
- 扩展 `get_headers()` 方法支持 Bearer Token 认证
- 兼容现有的 x-goog-api-key 认证方式
- **增强 Gemini 客户端** (`clients/gemini_client.py`)
- 新增 `query_balance_async()` 异步查询余额方法
- 新增 `query_balance_sync()` 同步查询余额方法(用于节点)
- 新增 `format_balance_info()` 格式化余额信息方法
- 余额转换公式:实际显示 = total_available / 500000
---
## [1.5.3] - 2026-02-01
### Changed
- **重大更新新手使用指南** (`GUIDE.md`)
- 新增批量提示词功能详细说明(节点详细说明部分)
- 新增场景 5:批量提示词生成(多提示词并发)
- 新增场景 6:批量提示词 + 图生图(共享参考图)
- 新增场景 7:批量提示词精准匹配规则详解
- 单提示词模式与批量提示词模式对比
- 详细的行为示例(纯文生图、图生图、多参考图)
- 常见误区与正确用法对照
- 使用决策树帮助用户选择合适的节点和模式
- 新增节点功能对比表和提示词模式对比表
- 新增 Q6-Q7:批量提示词相关常见问题
- 新增技巧 7:批量提示词最佳实践(4个子技巧)
- 更新场景编号(原场景 6-8 → 新场景 8)
- 修正批量提示词的描述(所有提示词共享输入图像,而非 1:1 匹配)
- 强化对 `---` 分隔符格式要求的说明
---
## [1.5.2] - 2026-02-01
### Added
- **新增新手使用指南** (`GUIDE.md`)
- 详细的安装配置步骤
- 三个核心节点的完整文档和参数说明
- 六种常见使用场景的工作流示例
- 常见问题解答和解决方案
- 六个进阶使用技巧
- 面向初次使用插件的用户,提供从零到一的完整指导
### Changed
- **更新开发规范** (`.cursorrules`)
- 新增"新手指南维护规则"章节
- 规定每次新增或变更节点时必须同步更新 `GUIDE.md`
- 提供节点文档和使用场景的标准模板
- 明确文档维护的五大原则(用户视角、实用性、完整性、同步性、可读性)
---
## [1.5.1] - 2026-02-01
### Fixed
- **修复批量 Nano Banana Pro 节点事件循环冲突**
- 修复 `RuntimeError: Cannot run the event loop while another loop is running` 错误
- 使用 `ThreadPoolExecutor` 在独立线程中运行异步事件循环
- 避免与 ComfyUI 主事件循环冲突
- 确保批量处理任务稳定执行
---
## [1.5.0] - 2026-02-01
### Changed
- **批量 Nano Banana Pro 节点重构** (`BatchNanoBananaPro`)
- 参数重命名:
- `手动参考图``加载参考图`
- `预览图像``输出图像`
- `配对模式``图片配对模式`
- 配对模式选项值重命名:
- `1:1索引配对``1:1`
- `笛卡尔积``1*N`
- 参数顺序调整:`文件夹2-4` 移到 `文件夹1` 下方(从 optional 移至 required
- 固定并发控制:移除 `最大并发数` 参数,默认自动管理(每批最多 100 并发)
- 固定生成数量:移除 `每组生成数量` 参数,每组配对固定生成 1 张图
### Added
- **批量 Nano Banana Pro 节点新增像素缩放功能**
- 新增 `像素缩放` 参数(BOOLEAN,默认 True
- 新增 `分辨率像素` 参数(FLOAT,默认 1.0,范围 0.1-100.0
- 支持对文件夹加载的图片和手动参考图进行缩放
- 使用 Lanczos 重采样算法保持图片质量
- 缩放在发送 API 前应用
### Removed
- **批量 Nano Banana Pro 节点移除处理报告功能**
- 移除返回值中的 `处理报告` 字符串输出
- 简化返回类型为单一 `IMAGE` 输出
- 统计信息仍通过控制台打印输出
### Improved
- **输出图像优化**
- `输出图像` 现在返回所有生成的图片(而非仅预览最后几张)
- 提供完整的批处理结果输出
---
## [1.4.0] - 2026-02-01
### Added
- **批量 Nano Banana Pro 节点** (`BatchNanoBananaPro`)
- 支持从 1-4 个文件夹批量加载图片
- 两种配对模式:1:1 索引配对 / 笛卡尔积
- 支持手动参考图输入(可与文件夹图片混合使用)
- 智能命名保存(保留原始文件名 + 自动后缀)
- 并发控制(默认最大 100,超过自动分批)
- 完整的处理报告输出
- 预览最后生成的图片
- **文件处理工具模块** (`utils/file_utils.py`)
- `load_images_from_folder()` - 从文件夹加载图片
- `pair_images_indexed()` - 1:1 索引配对
- `pair_images_cartesian()` - 笛卡尔积配对
- `generate_output_filename()` - 智能输出文件名生成
- `save_image()` - 保存图片到指定路径
### Technical
- 新增 `ImageInfo` 命名元组,携带图片元数据
- 支持 jpg/jpeg/png/webp/bmp/gif 图片格式
- 文件名按字母顺序排序,确保配对顺序一致
---
## [1.3.0] - 2026-02-01
### Added
- **批量提示词功能(Nano Banana Pro**
- 支持使用单行 `---` 分隔符同时提交多个不同提示词
- 所有提示词并发生成,提高效率
- 每个提示词可生成指定数量的图像(提示词数量 × 生图数量)
- 示例:3个提示词 × 2张/提示词 = 6张图
- 触发条件:`---` 必须单独占据一行
- 自动过滤空提示词
### Changed
- 优化 Nano Banana Pro 节点日志输出
- 批量提示词模式显示 "X个提示词 × Y张/提示词 = Z张图"
- 单提示词模式保持原有输出格式
### Technical
- 新增 `parse_batch_prompts()` 工具函数(utils/image_utils.py
- 新增 `generate_multi_prompts_async()` 方法(clients/gemini_client.py
- 新增 `generate_multi_prompts_sync()` 方法(clients/gemini_client.py
---
## [1.2.1] - 2026-02-01
### Changed
- **加载批次图像(Nano Banana Pro** 节点优化
- 节点显示名称改为 "加载批次图像(Nano Banana Pro"
- 移除 "参考图数量" 参数,改为自动检测所有有效输入图像数量
- "开启像素缩放" 改名为 "像素缩放",类型改为 BOOLEAN 开关(默认开启)
- "目标像素数_百万" 改名为 "分辨率像素"
### Improved
- 简化用户操作流程,无需手动设置图像数量
- 支持灵活的图像输入方式(1-14 张任意数量)
---
## [1.2.0] - 2026-02-01
### Added
- **加载批次图像** 节点 (`BatchImageLoader`)
- 支持动态输入数量(1-14 张图像)
- 可选的像素缩放功能(保持纵横比)
- 使用 Lanczos 重采样算法进行高质量缩放
- 支持设置目标像素数(0.1-100 百万像素)
- 可与原生"加载图像"节点连接使用
- 输出批次张量供其他节点使用
---
## [1.1.1] - 2026-02-01
### Changed
- 将 Nano Banana Pro 节点的随机种子参数改为 ComfyUI 原生格式
- `随机种子``seed` (参数名符合 ComfyUI 标准)
- 移除 `-1` 自动随机逻辑
- 种子默认值改为 `0`,取值范围 `[0, 2^64-1]`
### Removed
- 移除 `control_after_generate` 参数(简化节点参数)
---
## [1.1.0] - 2026-02-01
### Changed
- 重构项目结构,模块化设计
- 新增 `nodes/` 目录存放节点实现
- 新增 `utils/` 目录存放工具函数
- 新增 `clients/` 目录存放 API 客户端
- 封装图像转换工具到 `utils/image_utils.py`
- 封装配置管理到 `utils/config.py`
- 重构 API 客户端,拆分为基类和具体实现
- 精简文档结构,整合为 README.md + CHANGELOG.md + .cursorrules
### Added
- `BaseAPIClient` 抽象基类,支持快速开发新 API 客户端
- `.cursorrules` 开发指导文档,统一开发规范
- `get_api_key_or_raise()` 函数,简化密钥获取逻辑
---
## [1.0.0] - 2026-02-01
### Added
- 初始版本发布
- Nano Banana Pro 节点
- 文生图功能
- 图生图功能(最多 14 张输入图像)
- 批量并发生成(最多 1000 张)
- 多分辨率支持(1K / 2K / 4K
- 10 种宽高比选项
- 可控随机种子
- 通过 api.o1key.com 调用 Gemini 3 Pro 模型
- 支持 .config 文件和环境变量配置 API 密钥
- 完善的错误处理和日志输出
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2026 Comfyui_o1key
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+12
View File
@@ -0,0 +1,12 @@
# Comfyui_o1key
通过 `api.o1key.com` 调用 AI 模型的 ComfyUI 自定义节点集合。
## 功能特性
- 🎨 文生图 / 图生图
- 🔄 批量并发生成(最多 1000 张)
- 📐 10 种宽高比
- 🎯 3 种分辨率(1K / 2K / 4K
- 🌱 可控随机种子
+37
View File
@@ -0,0 +1,37 @@
"""
Comfyui_o1key - ComfyUI 自定义节点集合
通过 api.o1key.com 调用 AI 模型进行图像生成和文本生成
项目结构:
├── nodes/ # 节点实现
├── utils/ # 工具模块
├── clients/ # API 客户端
└── __init__.py # 节点注册入口
"""
# 检查更新(仅在启动时检查一次)
try:
from .utils.update_checker import check_for_updates, notify_update_available
if check_for_updates():
notify_update_available()
except Exception:
# 静默失败,不影响插件加载
pass
from .nodes import NanoBananaPro, BatchNanoBananaPro, GoogleGemini
# ComfyUI 节点注册
NODE_CLASS_MAPPINGS = {
"NanoBananaPro": NanoBananaPro,
"BatchNanoBananaPro": BatchNanoBananaPro,
"GoogleGemini": GoogleGemini
}
NODE_DISPLAY_NAME_MAPPINGS = {
"NanoBananaPro": "Nano Banana Pro",
"BatchNanoBananaPro": "批量 Nano Banana Pro",
"GoogleGemini": "Google Gemini"
}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+10
View File
@@ -0,0 +1,10 @@
"""
API 客户端模块
包含与外部 API 通信的客户端实现
"""
from .base_client import BaseAPIClient
from .gemini_client import GeminiAPIClient
from .gemini_flash_client import GeminiFlashClient
__all__ = ['BaseAPIClient', 'GeminiAPIClient', 'GeminiFlashClient']
+368
View File
@@ -0,0 +1,368 @@
"""
API 客户端基类
提供通用的 HTTP 请求、响应解析和错误处理功能
"""
import asyncio
import json
import threading
from abc import ABC, abstractmethod
from typing import Any, Callable, Dict, List, Optional
import aiohttp
class BaseAPIClient(ABC):
"""
API 客户端抽象基类
子类需要实现以下方法:
- get_endpoint(): 获取 API 端点
- build_request_body(): 构建请求体
- parse_response(): 解析响应
"""
def __init__(
self,
base_url: str,
api_key: str,
max_request_size: int = 20 * 1024 * 1024
):
"""
初始化客户端
Args:
base_url: API 基础 URL
api_key: API 密钥
max_request_size: 最大请求体大小(字节),默认 20MB
"""
self.base_url = base_url
self.api_key = api_key
self.max_request_size = max_request_size
@abstractmethod
def get_endpoint(self, **kwargs) -> str:
"""
获取 API 端点路径
Args:
**kwargs: 额外参数(如模型名、分辨率等)
Returns:
端点路径字符串
"""
pass
@abstractmethod
def build_request_body(self, **kwargs) -> Dict[str, Any]:
"""
构建 API 请求体
Args:
**kwargs: 请求参数
Returns:
请求体字典
"""
pass
@abstractmethod
def parse_response(self, response: Dict[str, Any]) -> Any:
"""
解析 API 响应
Args:
response: API 响应字典
Returns:
解析后的结果
"""
pass
def get_headers(self, use_bearer_token: bool = False) -> Dict[str, str]:
"""
获取请求头
Args:
use_bearer_token: 是否使用 Bearer Token 认证(默认为 False
Returns:
请求头字典
"""
if use_bearer_token:
return {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json"
}
else:
return {
"x-goog-api-key": self.api_key,
"Content-Type": "application/json"
}
def check_request_size(self, request_body: Dict[str, Any]) -> None:
"""
检查请求体大小是否超过限制
Args:
request_body: 请求体字典
Raises:
ValueError: 如果请求体超过限制
"""
request_json = json.dumps(request_body)
request_size = len(request_json.encode('utf-8'))
if request_size > self.max_request_size:
size_mb = request_size / 1024 / 1024
limit_mb = self.max_request_size / 1024 / 1024
raise ValueError(
f"请求体大小 {size_mb:.2f}MB 超过限制 {limit_mb:.0f}MB"
"请降低分辨率或减少图片数量"
)
async def request_async(
self,
endpoint: str,
request_body: Dict[str, Any],
session: Optional[aiohttp.ClientSession] = None,
use_bearer_token: bool = False,
timeout: Optional[int] = None
) -> Dict[str, Any]:
"""
发送异步 HTTP 请求
Args:
endpoint: API 端点
request_body: 请求体
session: aiohttp 会话(可选)
use_bearer_token: 是否使用 Bearer Token 认证
Returns:
响应 JSON
Raises:
RuntimeError: 请求失败时
"""
url = f"{self.base_url}{endpoint}"
headers = self.get_headers(use_bearer_token)
# 检查请求大小
self.check_request_size(request_body)
close_session = False
if session is None:
session = aiohttp.ClientSession()
close_session = True
try:
# 设置超时
timeout_obj = aiohttp.ClientTimeout(total=timeout) if timeout else None
async with session.post(url, json=request_body, headers=headers, timeout=timeout_obj) as response:
if response.status != 200:
error_text = await response.text()
# 针对常见错误状态码提供友好提示
if response.status == 504:
raise RuntimeError(
f"API 请求超时 (504 Gateway Timeout)\n"
f"原因:服务器响应超时或该端点暂时不可用\n"
f"建议:\n"
f" - 尝试使用其他模型\n"
f" - 稍后重试\n"
f" - 降低分辨率或减少输入图像数量\n"
f"详细错误: {error_text[:200]}"
)
elif response.status == 503:
raise RuntimeError(
f"服务暂时不可用 (503 Service Unavailable)\n"
f"原因:模型服务过载或维护中\n"
f"建议:\n"
f" - 稍后重试\n"
f" - 尝试使用其他模型"
)
elif response.status == 429:
raise RuntimeError(
f"请求频率超限 (429 Too Many Requests)\n"
f"原因:API 配额用尽或请求过于频繁\n"
f"建议:\n"
f" - 等待一段时间后重试\n"
f" - 检查 API 配额是否充足"
)
elif response.status == 404:
raise RuntimeError(
f"端点不存在 (404 Not Found)\n"
f"原因:API 端点路径错误或模型不存在\n"
f"建议:\n"
f" - 检查模型名称是否正确\n"
f" - 使用其他可用模型"
)
else:
raise RuntimeError(
f"API 请求失败 (状态码: {response.status}): {error_text}"
)
return await response.json()
finally:
if close_session:
await session.close()
async def request_get_async(
self,
endpoint: str,
session: Optional[aiohttp.ClientSession] = None,
use_bearer_token: bool = True,
timeout: Optional[int] = None
) -> Dict[str, Any]:
"""
发送异步 HTTP GET 请求
Args:
endpoint: API 端点
session: aiohttp 会话(可选)
use_bearer_token: 是否使用 Bearer Token 认证(默认为 True
Returns:
响应 JSON
Raises:
RuntimeError: 请求失败时
"""
url = f"{self.base_url}{endpoint}"
headers = self.get_headers(use_bearer_token)
close_session = False
if session is None:
session = aiohttp.ClientSession()
close_session = True
try:
# 设置超时
timeout_obj = aiohttp.ClientTimeout(total=timeout) if timeout else None
async with session.get(url, headers=headers, timeout=timeout_obj) as response:
if response.status != 200:
error_text = await response.text()
# 针对常见错误状态码提供友好提示
if response.status == 504:
raise RuntimeError(
f"API 请求超时 (504 Gateway Timeout)\n"
f"原因:服务器响应超时或该端点暂时不可用\n"
f"建议:稍后重试"
)
elif response.status == 503:
raise RuntimeError(
f"服务暂时不可用 (503 Service Unavailable)\n"
f"原因:服务过载或维护中\n"
f"建议:稍后重试"
)
elif response.status == 429:
raise RuntimeError(
f"请求频率超限 (429 Too Many Requests)\n"
f"原因:API 配额用尽或请求过于频繁\n"
f"建议:等待一段时间后重试"
)
else:
raise RuntimeError(
f"API 请求失败 (状态码: {response.status}): {error_text}"
)
return await response.json()
finally:
if close_session:
await session.close()
async def batch_request_async(
self,
requests: List[Dict[str, Any]],
progress_callback: Optional[Callable[[int, int], None]] = None
) -> List[Any]:
"""
批量并发请求
Args:
requests: 请求列表,每个元素包含 endpoint 和 request_body
progress_callback: 进度回调函数 (current, total)
Returns:
响应结果列表
"""
results = []
completed = 0
total = len(requests)
# 创建无限制的连接器
connector = aiohttp.TCPConnector(limit=0, limit_per_host=0)
async with aiohttp.ClientSession(connector=connector) as session:
tasks = []
for req in requests:
task = self.request_async(
endpoint=req['endpoint'],
request_body=req['request_body'],
session=session
)
tasks.append(task)
# 并发执行
responses = await asyncio.gather(*tasks, return_exceptions=True)
for i, resp in enumerate(responses):
if isinstance(resp, Exception):
print(f"⚠️ 第 {i+1} 个请求失败: {str(resp)}")
continue
try:
parsed = self.parse_response(resp)
results.append(parsed)
completed += 1
if progress_callback:
progress_callback(completed, total)
except Exception as e:
print(f"⚠️ 第 {i+1} 个响应解析失败: {str(e)}")
return results
def run_async_in_thread(self, coro) -> Any:
"""
在独立线程中运行异步代码(用于 ComfyUI 同步接口)
Args:
coro: 协程对象
Returns:
协程执行结果
"""
result_container = []
error_container = []
def run_in_thread():
try:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
result = loop.run_until_complete(coro)
result_container.append(result)
finally:
loop.close()
except Exception as e:
error_container.append(e)
thread = threading.Thread(target=run_in_thread)
thread.start()
thread.join()
if error_container:
raise error_container[0]
if not result_container:
raise RuntimeError("异步任务未返回结果")
return result_container[0]
+728
View File
@@ -0,0 +1,728 @@
"""
Gemini API 客户端
处理与 api.o1key.com 的通信,用于图像生成
"""
import re
import time
from io import BytesIO
from typing import Any, Callable, Dict, List, Optional
import aiohttp
from PIL import Image
from ..utils.image_utils import encode_image_to_base64, decode_base64_to_pil
from ..utils.config import get_api_key_or_raise
from .base_client import BaseAPIClient
# API 基础配置
API_BASE_URL = "https://api.o1key.com"
class GeminiAPIClient(BaseAPIClient):
"""
Gemini API 客户端
用于调用 Gemini 3 Pro 模型进行图像生成
"""
@staticmethod
def get_timeout_by_resolution(resolution: str) -> int:
"""
根据分辨率获取超时时间
Args:
resolution: 分辨率(1K, 2K, 4K
Returns:
超时时间(秒)
"""
timeout_map = {
"1K": 180, # 3 分钟
"2K": 300, # 5 分钟
"4K": 360 # 6 分钟
}
return timeout_map.get(resolution, 300) # 默认 5 分钟
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=API_BASE_URL,
api_key=api_key,
max_request_size=20 * 1024 * 1024
)
def get_endpoint(self, model: str = "", resolution: str = "2K", **kwargs) -> str:
"""
根据模型和分辨率获取 API 端点
Args:
model: 模型名称
resolution: 分辨率(1K, 2K, 4K
Returns:
API 端点路径
"""
from ..models_config import get_model_endpoint
# 特殊处理:动态端点模型(根据分辨率选择)
if model == "nano-banana-pro":
if resolution == "1K":
return "/v1beta/models/nano-banana-pro:generateContent"
elif resolution == "2K":
return "/v1beta/models/nano-banana-pro-2k:generateContent"
elif resolution == "4K":
return "/v1beta/models/nano-banana-pro-4k:generateContent"
else:
return "/v1beta/models/nano-banana-pro-2k:generateContent"
elif model == "gemini-3-pro-image-preview-url":
if resolution == "1K":
return "/v1beta/models/gemini-3-pro-image-preview-url:generateContent"
elif resolution == "2K":
return "/v1beta/models/gemini-3-pro-image-preview-2k-url:generateContent"
elif resolution == "4K":
return "/v1beta/models/gemini-3-pro-image-preview-4k-url:generateContent"
else:
return "/v1beta/models/gemini-3-pro-image-preview-2k-url:generateContent"
# 其他模型:从配置文件读取端点
endpoint = get_model_endpoint(model)
if endpoint:
return endpoint
# 兜底:使用标准模式端点
return "/v1beta/models/gemini-3-pro-image-preview:generateContent"
def build_request_body(
self,
prompt: str = "",
images: Optional[List[Image.Image]] = None,
aspect_ratio: str = "1:1",
resolution: str = "2K",
**kwargs
) -> Dict[str, Any]:
"""
构建 API 请求体
Args:
prompt: 提示词
images: 输入图像列表(可选)
aspect_ratio: 宽高比
resolution: 分辨率
Returns:
请求体字典
"""
parts = []
# 添加文本部分
parts.append({"text": prompt})
# 添加图像部分(如果有)
if images:
for img in images:
img_base64 = encode_image_to_base64(img)
parts.append({
"inline_data": {
"mime_type": "image/png",
"data": img_base64
}
})
# 构建请求体
request_body = {
"contents": [
{
"role": "user",
"parts": parts
}
],
"generationConfig": {
"responseModalities": ["TEXT", "IMAGE"],
"imageConfig": {
"aspectRatio": aspect_ratio,
"imageSize": resolution
}
}
}
return request_body
def parse_response(self, response: Dict[str, Any]) -> List[Image.Image]:
"""
同步解析 API 响应(保留以满足抽象基类要求)
注意:此方法仅用于兼容基类接口,实际使用请调用 parse_response_async()
Args:
response: API 响应字典
Returns:
图像列表
Raises:
RuntimeError: 此方法不应被直接调用
"""
raise RuntimeError(
"parse_response() 不应被直接调用。"
"请使用 generate_single_async() 或 generate_batch_async() 等高级方法。"
)
async def parse_response_async(
self,
response: Dict[str, Any],
session: Optional[aiohttp.ClientSession] = None
) -> List[Image.Image]:
"""
异步解析 API 响应,提取生成的图像
Args:
response: API 响应字典
session: aiohttp 会话(用于下载图片)
Returns:
图像列表
Raises:
RuntimeError: 解析失败或 API 拒绝时
"""
# ========== 错误检测(按优先级顺序)==========
# 1. 检查 candidatesTokenCount(最高优先级)
usage_metadata = response.get("usageMetadata", {})
candidates_token_count = usage_metadata.get("candidatesTokenCount", -1)
if candidates_token_count == 0:
error_msg = (
"内容审核拒绝 - candidatesTokenCount = 0\n\n"
"原因:提示词或参考图包含不适当内容(色情、暴力、敏感话题等),"
"在内容审核阶段就被拒绝,连候选内容都未生成。\n\n"
"建议:\n"
" - 检查提示词,确保不包含敏感或违规内容\n"
" - 如使用参考图,确保图片内容健康合规\n"
" - 避免描述暴力、色情等不当内容\n"
" - 调整提示词后重试"
)
raise RuntimeError(error_msg)
# 2. 检查 finishReason(次优先级)
candidates = response.get("candidates", [])
if candidates:
for candidate in candidates:
finish_reason = candidate.get("finishReason", "")
if finish_reason and finish_reason != "STOP":
# 根据不同的 finishReason 提供具体建议
reason_messages = {
"PROHIBITED_CONTENT": (
"违禁内容拒绝",
"生成内容触发了违禁内容策略",
[
"避免引用未来未发布的产品或概念(知识库截止2025年1月)",
"使用专业图片编辑软件处理特殊需求",
"确保请求内容在模型知识范围内"
]
),
"SAFETY": (
"安全过滤器拒绝",
"内容触发了安全过滤器",
[
"使用健康、正面的描述",
"避免涉及隐私和伦理问题的内容",
"调整提示词后重试"
]
),
"RECITATION": (
"版权问题",
"可能涉及版权或重复已有内容",
[
"避免涉及版权敏感话题",
"使用更原创的描述方式",
"调整提示词后重试"
]
),
"MAX_TOKENS": (
"Token 超限",
"生成的内容超过了 Token 限制",
[
"简化提示词",
"减少输入图片数量",
"降低请求复杂度"
]
)
}
if finish_reason in reason_messages:
title, reason, suggestions = reason_messages[finish_reason]
suggestions_text = "\n".join([f" - {s}" for s in suggestions])
error_msg = (
f"{title} - finishReason = {finish_reason}\n\n"
f"原因:{reason}\n\n"
f"建议:\n{suggestions_text}"
)
else:
# 未知的 finishReason
error_msg = (
f"生成异常 - finishReason = {finish_reason}\n\n"
"原因:生成过程中断,具体原因未知\n\n"
"建议:\n"
" - 使用健康、正面的描述\n"
" - 避免敏感话题\n"
" - 调整提示词后重试"
)
raise RuntimeError(error_msg)
# ========== 图像提取 ==========
images = []
text_responses = [] # 收集文本响应
# 需要关闭 session 的标记
close_session = False
if session is None:
session = aiohttp.ClientSession()
close_session = True
try:
for candidate in candidates:
content = candidate.get("content", {})
parts = content.get("parts", [])
for part in parts:
# 方式1: inline_data 或 inlineData (base64)
# 兼容两种命名方式:蛇形(inline_data)和驼峰(inlineData
inline_data_key = None
if "inline_data" in part:
inline_data_key = "inline_data"
elif "inlineData" in part:
inline_data_key = "inlineData"
if inline_data_key:
inline_data = part[inline_data_key]
# 同样兼容 data/mimeType 的命名
img_data = inline_data.get("data") or inline_data.get("data", "")
if img_data:
img = decode_base64_to_pil(img_data)
images.append(img)
# 方式2: text 中的 URL - 改为异步下载
elif "text" in part:
text = part["text"]
# 收集文本响应(用于后续错误检测)
text_responses.append(text)
# 尝试 markdown 格式: ![alt](url)
url_pattern_md = r'!\[.*?\]\((https?://[^\)]+)\)'
urls = re.findall(url_pattern_md, text)
# 如果没找到,尝试纯 URL 格式
if not urls:
url_pattern_plain = r'https?://[^\s<>"{}|\\^`\[\]]+'
urls = re.findall(url_pattern_plain, text)
if urls:
for url in urls:
try:
# 使用 aiohttp 异步下载,支持更大的超时
download_start = time.time()
timeout = aiohttp.ClientTimeout(total=120)
async with session.get(url, timeout=timeout) as img_response:
if img_response.status == 200:
img_data = await img_response.read()
download_time = time.time() - download_start
img_size_mb = len(img_data) / 1024 / 1024
speed_mbps = img_size_mb / download_time if download_time > 0 else 0
# print(f"🔽 图片下载: {img_size_mb:.2f}MB 耗时 {download_time:.2f}s 速度 {speed_mbps:.2f}MB/s")
img = Image.open(BytesIO(img_data))
images.append(img)
else:
print(f"Nano Banana Pro: 下载图片失败 - HTTP {img_response.status}")
except Exception as e:
print(f"Nano Banana Pro: 下载图片失败 - {str(e)}")
# 方式3: 直接的 URL 字段 - 也改为异步
elif "imageUrl" in part or "url" in part:
url = part.get("imageUrl") or part.get("url")
try:
download_start = time.time()
timeout = aiohttp.ClientTimeout(total=120)
async with session.get(url, timeout=timeout) as img_response:
if img_response.status == 200:
img_data = await img_response.read()
download_time = time.time() - download_start
img_size_mb = len(img_data) / 1024 / 1024
speed_mbps = img_size_mb / download_time if download_time > 0 else 0
# print(f"🔽 图片下载: {img_size_mb:.2f}MB 耗时 {download_time:.2f}s 速度 {speed_mbps:.2f}MB/s")
img = Image.open(BytesIO(img_data))
images.append(img)
else:
print(f"Nano Banana Pro: 下载图片失败 - HTTP {img_response.status}")
except Exception as e:
print(f"Nano Banana Pro: 下载图片失败 - {str(e)}")
except Exception as e:
raise RuntimeError(f"解析 API 响应失败: {str(e)}")
finally:
if close_session:
await session.close()
# 3. 检查 API 文本响应拒绝说明
if not images and text_responses:
# API 返回了文本但没有图片,说明请求被拒绝
combined_text = "\n".join(text_responses)
error_msg = (
f"API 拒绝响应\n\n"
f"API 返回说明:\n{combined_text}\n\n"
f"建议:\n"
f" - 根据上述说明调整请求内容\n"
f" - 确保提示词和参考图符合使用规范"
)
raise RuntimeError(error_msg)
if not images:
raise RuntimeError("API 响应中未找到生成的图像")
return images
async def generate_single_async(
self,
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
images: Optional[List[Image.Image]] = None,
session=None
) -> List[Image.Image]:
"""
单次异步生成请求
Args:
prompt: 提示词
model: 模型名称
resolution: 分辨率
aspect_ratio: 宽高比
images: 输入图像列表
session: aiohttp 会话
Returns:
生成的图像列表
"""
endpoint = self.get_endpoint(model=model, resolution=resolution)
request_body = self.build_request_body(
prompt=prompt,
images=images,
aspect_ratio=aspect_ratio,
resolution=resolution
)
# 根据分辨率获取超时时间
timeout = self.get_timeout_by_resolution(resolution)
response = await self.request_async(endpoint, request_body, session, timeout=timeout)
# 使用异步解析方法,传入 session 以实现并发图片下载
return await self.parse_response_async(response, session)
async def generate_batch_async(
self,
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
batch_size: int,
images: Optional[List[Image.Image]] = None,
progress_callback: Optional[Callable[[int, int, bool, Optional[str]], None]] = None
) -> List[Image.Image]:
"""
批量全并发生成
Args:
prompt: 提示词
model: 模型名称
resolution: 分辨率
aspect_ratio: 宽高比
batch_size: 批次大小
images: 输入图像列表
progress_callback: 进度回调,签名为 (completed, total, success, error_msg)
Returns:
生成的图像列表
"""
import aiohttp
import asyncio
all_images = []
completed = 0
success_count = 0
fail_count = 0
connector = aiohttp.TCPConnector(limit=0, limit_per_host=0)
async with aiohttp.ClientSession(connector=connector) as session:
tasks = []
for i in range(batch_size):
task = asyncio.create_task(
self.generate_single_async(
prompt=prompt,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
images=images,
session=session
),
name=f"task_{i}"
)
tasks.append(task)
# 使用 as_completed 实时获取完成的任务
for coro in asyncio.as_completed(tasks):
completed += 1
try:
result = await coro
if result:
all_images.append(result[0])
success_count += 1
if progress_callback:
progress_callback(completed, batch_size, True, None)
except Exception as e:
fail_count += 1
error_msg = str(e)
# 截取错误信息的第一行
if '\n' in error_msg:
error_msg = error_msg.split('\n')[0]
if progress_callback:
progress_callback(completed, batch_size, False, error_msg)
if not all_images:
raise RuntimeError(f"批量生成失败,{fail_count} 个请求全部失败")
return all_images
def generate_sync(
self,
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
batch_size: int,
images: Optional[List[Image.Image]] = None,
progress_callback: Optional[Callable[[int, int], None]] = None
) -> List[Image.Image]:
"""
同步生成接口(用于 ComfyUI)
Args:
prompt: 提示词
model: 模型名称
resolution: 分辨率
aspect_ratio: 宽高比
batch_size: 批次大小
images: 输入图像列表
progress_callback: 进度回调
Returns:
生成的图像列表
"""
coro = self.generate_batch_async(
prompt=prompt,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
batch_size=batch_size,
images=images,
progress_callback=progress_callback
)
return self.run_async_in_thread(coro)
async def generate_multi_prompts_async(
self,
prompts: List[str],
model: str,
resolution: str,
aspect_ratio: str,
images_per_prompt: int,
images: Optional[List[Image.Image]] = None,
progress_callback: Optional[Callable[[int, int, bool, Optional[str]], None]] = None
) -> List[Image.Image]:
"""
多提示词批量生成
为每个提示词生成指定数量的图像,所有请求并发执行。
Args:
prompts: 提示词列表
model: 模型名称
resolution: 分辨率
aspect_ratio: 宽高比
images_per_prompt: 每个提示词生成的图像数量
images: 输入图像列表(所有提示词共享)
progress_callback: 进度回调,签名为 (completed, total, success, error_msg)
Returns:
生成的图像列表(长度 = len(prompts) * images_per_prompt
"""
import aiohttp
import asyncio
all_images = []
completed = 0
success_count = 0
fail_count = 0
total_tasks = len(prompts) * images_per_prompt
connector = aiohttp.TCPConnector(limit=0, limit_per_host=0)
async with aiohttp.ClientSession(connector=connector) as session:
tasks = []
# 为每个提示词创建 images_per_prompt 个任务
task_idx = 0
for prompt in prompts:
for _ in range(images_per_prompt):
task = asyncio.create_task(
self.generate_single_async(
prompt=prompt,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
images=images,
session=session
),
name=f"task_{task_idx}"
)
tasks.append(task)
task_idx += 1
# 使用 as_completed 实时获取完成的任务
for coro in asyncio.as_completed(tasks):
completed += 1
try:
result = await coro
if result:
all_images.append(result[0])
success_count += 1
if progress_callback:
progress_callback(completed, total_tasks, True, None)
except Exception as e:
fail_count += 1
error_msg = str(e)
# 截取错误信息的第一行
if '\n' in error_msg:
error_msg = error_msg.split('\n')[0]
if progress_callback:
progress_callback(completed, total_tasks, False, error_msg)
if not all_images:
raise RuntimeError(f"批量生成失败,{fail_count} 个请求全部失败")
return all_images
def generate_multi_prompts_sync(
self,
prompts: List[str],
model: str,
resolution: str,
aspect_ratio: str,
images_per_prompt: int,
images: Optional[List[Image.Image]] = None,
progress_callback: Optional[Callable[[int, int], None]] = None
) -> List[Image.Image]:
"""
多提示词批量生成(同步接口,用于 ComfyUI)
Args:
prompts: 提示词列表
model: 模型名称
resolution: 分辨率
aspect_ratio: 宽高比
images_per_prompt: 每个提示词生成的图像数量
images: 输入图像列表
progress_callback: 进度回调
Returns:
生成的图像列表
"""
coro = self.generate_multi_prompts_async(
prompts=prompts,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
images_per_prompt=images_per_prompt,
images=images,
progress_callback=progress_callback
)
return self.run_async_in_thread(coro)
async def query_balance_async(self) -> Dict[str, Any]:
"""
异步查询余额信息
Returns:
余额信息字典,包含:
- name: API 名称
- total_available: 可用余额(原始值)
- 其他字段...
Raises:
RuntimeError: 查询失败时
"""
endpoint = "/api/usage/token"
response = await self.request_get_async(endpoint, use_bearer_token=True)
if not response.get("code"):
raise RuntimeError("余额查询响应格式错误")
data = response.get("data", {})
return data
def query_balance_sync(self) -> Dict[str, Any]:
"""
同步查询余额信息(用于 ComfyUI 节点)
Returns:
余额信息字典
Raises:
RuntimeError: 查询失败时
"""
coro = self.query_balance_async()
return self.run_async_in_thread(coro)
def format_balance_info(self, balance_data: Dict[str, Any]) -> str:
"""
格式化余额信息为展示文本
Args:
balance_data: 余额信息字典
Returns:
格式化的文本,格式为 "当前余额:$XX.XX | APIxxx"
Example:
>>> data = {"name": "test-api", "total_available": 50000000}
>>> client.format_balance_info(data)
'当前余额:$100.00 | APItest-api'
"""
api_name = balance_data.get("name", "未知")
total_available = balance_data.get("total_available", 0)
# 转换公式:实际显示 = total_available / 500000
balance_in_dollars = total_available / 500000
return f"当前余额:${balance_in_dollars:.2f} | API{api_name}"
+273
View File
@@ -0,0 +1,273 @@
"""
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
from ..models_config import get_flash_model_endpoint, get_enabled_flash_models
from .base_client import BaseAPIClient
# API 基础配置
API_BASE_URL = "https://api.o1key.com"
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=API_BASE_URL,
api_key=api_key,
max_request_size=20 * 1024 * 1024 # 20MB
)
def get_endpoint(
self,
model: str = "gemini-3-flash-preview",
thinking_depth: str = "不思考",
**kwargs
) -> str:
"""
根据模型和思考深度获取 API 端点
Args:
model: 模型名称
thinking_depth: 思考深度 ("不思考""")
Returns:
API 端点路径
"""
endpoint = get_flash_model_endpoint(model, thinking_depth)
if endpoint is None:
# 回退到默认端点
default_models = get_enabled_flash_models()
if default_models:
endpoint = get_flash_model_endpoint(default_models[0], thinking_depth)
if endpoint is None:
raise ValueError(f"无法获取模型 '{model}' 的端点 (思考深度: {thinking_depth})")
return endpoint
def build_request_body(
self,
prompt: str = "",
system_instruction: Optional[str] = None,
image_data: Optional[List[Dict[str, str]]] = None,
video_data: Optional[Dict[str, str]] = None,
**kwargs
) -> Dict[str, Any]:
"""
构建 API 请求体
Args:
prompt: 用户提示词
system_instruction: 系统指令(可选)
image_data: 图片数据列表,每个元素包含 mime_type 和 data
video_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"]
}
})
# 构建请求体
request_body = {
"contents": [
{
"parts": parts
}
]
}
# 添加系统指令(如果有)
if system_instruction and system_instruction.strip():
request_body["system_instruction"] = {
"parts": [
{"text": system_instruction}
]
}
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_depth: str = "不思考",
system_instruction: Optional[str] = None,
image_data: Optional[List[Dict[str, str]]] = None,
video_data: Optional[Dict[str, str]] = None,
session: Optional[aiohttp.ClientSession] = None
) -> str:
"""
异步生成文本
Args:
prompt: 用户提示词
model: 模型名称
thinking_depth: 思考深度
system_instruction: 系统指令
image_data: 图片数据列表
video_data: 视频数据
session: aiohttp 会话
Returns:
生成的文本内容
"""
endpoint = self.get_endpoint(model=model, thinking_depth=thinking_depth)
request_body = self.build_request_body(
prompt=prompt,
system_instruction=system_instruction,
image_data=image_data,
video_data=video_data
)
# 根据是否有视频设置超时(视频处理需要更长时间)
timeout = 300 if video_data else 180
response = await self.request_async(
endpoint,
request_body,
session,
timeout=timeout
)
return self.parse_response(response)
def generate_sync(
self,
prompt: str,
model: str = "gemini-3-flash-preview",
thinking_depth: str = "不思考",
system_instruction: Optional[str] = None,
image_data: Optional[List[Dict[str, str]]] = None,
video_data: Optional[Dict[str, str]] = None
) -> str:
"""
同步生成文本(用于 ComfyUI 节点)
Args:
prompt: 用户提示词
model: 模型名称
thinking_depth: 思考深度
system_instruction: 系统指令
image_data: 图片数据列表
video_data: 视频数据
Returns:
生成的文本内容
"""
coro = self.generate_async(
prompt=prompt,
model=model,
thinking_depth=thinking_depth,
system_instruction=system_instruction,
image_data=image_data,
video_data=video_data
)
return self.run_async_in_thread(coro)
+449
View File
@@ -0,0 +1,449 @@
"""
模型配置中心
用于集中管理所有支持的 Gemini 模型
使用方式:
1. 添加新模型: 在对应的模型列表中添加新的模型字典
2. 临时关闭模型: 将模型的 enabled 字段设为 False
3. 重新启用模型: 将模型的 enabled 字段改回 True
模型类型:
- GEMINI_MODELS: Nano Banana Pro 图像生成模型
- GEMINI_FLASH_MODELS: Google Gemini Flash 文本生成模型
示例:
添加新模型:
{
"id": "gemini-新模型名称",
"description": "模型说明和特点",
"enabled": True,
"endpoint_type": "standard" # 端点类型: "dynamic", "standard", "flatfee"
}
临时关闭模型:
将对应模型的 "enabled": True 改为 "enabled": False
"""
from typing import List, Dict, Optional
# ============================================================
# 模型配置列表
# ============================================================
# ============================================================
# Nano Banana Pro 图像生成模型
# ============================================================
GEMINI_MODELS = [
{
"id": "nano-banana-pro",
"description": "Nano Banana Pro,根据分辨率自动选择端点 (1K/2K/4K),高性能图像生成模型",
"enabled": True,
"endpoint_type": "dynamic",
"endpoint": None # 动态端点,由代码根据分辨率选择
},
{
"id": "gemini-3-pro-image-preview-url",
"description": "URL 模式,根据分辨率自动选择端点 (1K/2K/4K),推荐用于需要不同分辨率的场景",
"enabled": False,
"endpoint_type": "dynamic",
"endpoint": None # 动态端点,由代码根据分辨率选择
},
{
"id": "gemini-3-pro-image-preview",
"description": "标准模式,固定端点,适用于常规图像生成",
"enabled": True,
"endpoint_type": "standard",
"endpoint": "/v1beta/models/gemini-3-pro-image-preview:generateContent"
},
{
"id": "gemini-3-pro-image-preview-flatfee",
"description": "固定费用模式,固定端点,按固定价格计费 (暂时不可用-504错误)",
"enabled": False, # 暂时禁用:端点返回 504 错误
"endpoint_type": "flatfee",
"endpoint": "/v1beta/models/gemini-3-pro-image-preview-flatfee:generateContent"
},
{
"id": "nano-banana-2",
"description": "Nano Banana 2 模型,固定端点,适用于高质量图像生成",
"enabled": False,
"endpoint_type": "standard",
"endpoint": "/v1beta/models/nano-banana-2:generateContent"
}
]
# ============================================================
# Google Gemini Flash 文本生成模型
# ============================================================
GEMINI_FLASH_MODELS = [
{
"id": "gemini-3-flash-preview",
"description": "Gemini 3 Flash,快速多模态文本生成,支持图片和视频输入",
"enabled": True,
"endpoints": {
"不思考": "/v1beta/models/gemini-3-flash-preview-nothinking:generateContent",
"": "/v1beta/models/gemini-3-flash-preview-high:generateContent"
}
}
]
# ============================================================
# 工具函数
# ============================================================
def get_enabled_models() -> List[str]:
"""
获取所有启用的模型 ID 列表
Returns:
启用的模型 ID 列表
Example:
>>> get_enabled_models()
['gemini-3-pro-image-preview-url', 'gemini-3-pro-image-preview', ...]
"""
return [model["id"] for model in GEMINI_MODELS if model.get("enabled", False)]
def get_all_models() -> List[str]:
"""
获取所有模型 ID 列表(包括已禁用的)
Returns:
所有模型 ID 列表
Example:
>>> get_all_models()
['gemini-3-pro-image-preview-url', 'gemini-3-pro-image-preview', ...]
"""
return [model["id"] for model in GEMINI_MODELS]
def get_model_config(model_id: str) -> Optional[Dict]:
"""
根据模型 ID 获取完整的模型配置
Args:
model_id: 模型 ID
Returns:
模型配置字典,如果未找到则返回 None
Example:
>>> config = get_model_config("gemini-3-pro-image-preview-url")
>>> print(config["description"])
URL 模式,根据分辨率自动选择端点 (1K/2K/4K)
"""
for model in GEMINI_MODELS:
if model["id"] == model_id:
return model
return None
def is_model_enabled(model_id: str) -> bool:
"""
检查指定模型是否启用
Args:
model_id: 模型 ID
Returns:
True 如果模型启用,False 如果禁用或不存在
Example:
>>> is_model_enabled("gemini-3-pro-image-preview-url")
True
"""
config = get_model_config(model_id)
if config is None:
return False
return config.get("enabled", False)
def get_model_description(model_id: str) -> str:
"""
获取模型的描述信息
Args:
model_id: 模型 ID
Returns:
模型描述,如果未找到则返回空字符串
Example:
>>> get_model_description("gemini-3-pro-image-preview")
'标准模式,固定端点,适用于常规图像生成'
"""
config = get_model_config(model_id)
if config is None:
return ""
return config.get("description", "")
def get_endpoint_type(model_id: str) -> Optional[str]:
"""
获取模型的端点类型
Args:
model_id: 模型 ID
Returns:
端点类型 ("dynamic", "standard", "flatfee"),如果未找到则返回 None
Example:
>>> get_endpoint_type("gemini-3-pro-image-preview-url")
'dynamic'
"""
config = get_model_config(model_id)
if config is None:
return None
return config.get("endpoint_type")
def get_model_endpoint(model_id: str) -> Optional[str]:
"""
获取模型的 API 端点
Args:
model_id: 模型 ID
Returns:
API 端点路径,如果未找到或为动态端点则返回 None
Example:
>>> get_model_endpoint("gemini-3-pro-image-preview")
'/v1beta/models/gemini-3-pro-image-preview:generateContent'
>>> get_model_endpoint("gemini-3-pro-image-preview-url")
None # 动态端点
"""
config = get_model_config(model_id)
if config is None:
return None
return config.get("endpoint")
# ============================================================
# Gemini Flash 模型工具函数
# ============================================================
def get_enabled_flash_models() -> List[str]:
"""
获取所有启用的 Flash 模型 ID 列表
Returns:
启用的 Flash 模型 ID 列表
Example:
>>> get_enabled_flash_models()
['gemini-3-flash-preview']
"""
return [model["id"] for model in GEMINI_FLASH_MODELS if model.get("enabled", False)]
def get_all_flash_models() -> List[str]:
"""
获取所有 Flash 模型 ID 列表(包括已禁用的)
Returns:
所有 Flash 模型 ID 列表
"""
return [model["id"] for model in GEMINI_FLASH_MODELS]
def get_flash_model_config(model_id: str) -> Optional[Dict]:
"""
根据模型 ID 获取 Flash 模型的完整配置
Args:
model_id: 模型 ID
Returns:
模型配置字典,如果未找到则返回 None
Example:
>>> config = get_flash_model_config("gemini-3-flash-preview")
>>> print(config["description"])
'Gemini 3 Flash,快速多模态文本生成,支持图片和视频输入'
"""
for model in GEMINI_FLASH_MODELS:
if model["id"] == model_id:
return model
return None
def is_flash_model_enabled(model_id: str) -> bool:
"""
检查指定 Flash 模型是否启用
Args:
model_id: 模型 ID
Returns:
True 如果模型启用,False 如果禁用或不存在
"""
config = get_flash_model_config(model_id)
if config is None:
return False
return config.get("enabled", False)
def get_flash_model_endpoint(model_id: str, thinking_depth: str = "不思考") -> Optional[str]:
"""
获取 Flash 模型的 API 端点
Args:
model_id: 模型 ID
thinking_depth: 思考深度 ("不思考""")
Returns:
API 端点路径,如果未找到则返回 None
Example:
>>> get_flash_model_endpoint("gemini-3-flash-preview", "不思考")
'/v1beta/models/gemini-3-flash-preview-nothinking:generateContent'
>>> get_flash_model_endpoint("gemini-3-flash-preview", "")
'/v1beta/models/gemini-3-flash-preview-high:generateContent'
"""
config = get_flash_model_config(model_id)
if config is None:
return None
endpoints = config.get("endpoints", {})
return endpoints.get(thinking_depth)
def get_flash_model_description(model_id: str) -> str:
"""
获取 Flash 模型的描述信息
Args:
model_id: 模型 ID
Returns:
模型描述,如果未找到则返回空字符串
"""
config = get_flash_model_config(model_id)
if config is None:
return ""
return config.get("description", "")
# ============================================================
# 向后兼容性检查
# ============================================================
def validate_models_config() -> None:
"""
验证模型配置的完整性
检查:
- 每个模型必须有 id, description, enabled, endpoint_type, endpoint 字段
- 非动态端点模型必须配置有效的 endpoint
- 至少有一个模型是启用的
Raises:
ValueError: 如果配置不合法
"""
if not GEMINI_MODELS:
raise ValueError("GEMINI_MODELS 列表不能为空")
required_fields = ["id", "description", "enabled", "endpoint_type", "endpoint"]
valid_endpoint_types = ["dynamic", "standard", "flatfee"]
for i, model in enumerate(GEMINI_MODELS):
# 检查必需字段
for field in required_fields:
if field not in model:
raise ValueError(f"模型 #{i} 缺少必需字段: {field}")
# 检查 endpoint_type 是否合法
if model["endpoint_type"] not in valid_endpoint_types:
raise ValueError(
f"模型 {model['id']} 的 endpoint_type '{model['endpoint_type']}' 不合法。"
f"必须是: {', '.join(valid_endpoint_types)}"
)
# 检查非动态端点模型必须有有效的 endpoint
if model["endpoint_type"] != "dynamic" and not model.get("endpoint"):
raise ValueError(
f"模型 {model['id']} 的 endpoint_type 为 '{model['endpoint_type']}'"
f"但未配置有效的 endpoint 字段"
)
# 检查端点格式(如果配置了)
endpoint = model.get("endpoint")
if endpoint and not endpoint.startswith("/v1beta/models/"):
raise ValueError(
f"模型 {model['id']} 的 endpoint '{endpoint}' 格式不正确。"
f"应以 '/v1beta/models/' 开头"
)
# 检查至少有一个启用的模型
if not get_enabled_models():
raise ValueError("至少需要启用一个模型")
def validate_flash_models_config() -> None:
"""
验证 Flash 模型配置的完整性
检查:
- 每个模型必须有 id, description, enabled, endpoints 字段
- endpoints 必须包含所有思考深度选项
- 至少有一个模型是启用的
Raises:
ValueError: 如果配置不合法
"""
if not GEMINI_FLASH_MODELS:
raise ValueError("GEMINI_FLASH_MODELS 列表不能为空")
required_fields = ["id", "description", "enabled", "endpoints"]
required_thinking_depths = ["不思考", ""]
for i, model in enumerate(GEMINI_FLASH_MODELS):
# 检查必需字段
for field in required_fields:
if field not in model:
raise ValueError(f"Flash 模型 #{i} 缺少必需字段: {field}")
# 检查 endpoints 字典
endpoints = model.get("endpoints", {})
if not isinstance(endpoints, dict):
raise ValueError(f"Flash 模型 {model['id']} 的 endpoints 必须是字典")
# 检查所有思考深度选项都有对应端点
for depth in required_thinking_depths:
if depth not in endpoints:
raise ValueError(
f"Flash 模型 {model['id']} 的 endpoints 缺少 '{depth}' 思考深度"
)
endpoint = endpoints[depth]
if not endpoint or not endpoint.startswith("/v1beta/models/"):
raise ValueError(
f"Flash 模型 {model['id']} 的端点 '{endpoint}' 格式不正确。"
f"应以 '/v1beta/models/' 开头"
)
# 检查至少有一个启用的模型
if not get_enabled_flash_models():
raise ValueError("至少需要启用一个 Flash 模型")
# 在模块加载时验证配置
try:
validate_models_config()
except ValueError as e:
print(f"⚠️ 图像模型配置验证失败: {str(e)}")
print(f"⚠️ 请检查 models_config.py 文件")
try:
validate_flash_models_config()
except ValueError as e:
print(f"⚠️ Flash 模型配置验证失败: {str(e)}")
print(f"⚠️ 请检查 models_config.py 文件")
+10
View File
@@ -0,0 +1,10 @@
"""
节点模块
包含所有 ComfyUI 自定义节点的实现
"""
from .nano_banana_pro import NanoBananaPro
from .batch_nano_banana_pro import BatchNanoBananaPro
from .google_gemini import GoogleGemini
__all__ = ['NanoBananaPro', 'BatchNanoBananaPro', 'GoogleGemini']
+751
View File
@@ -0,0 +1,751 @@
"""
批量 Nano Banana Pro 节点
ComfyUI 自定义节点,用于批量处理图像生成任务
支持多文件夹加载、1:1/笛卡尔积配对、智能命名保存
"""
import time
import math
import random
import asyncio
import aiohttp
from concurrent.futures import ThreadPoolExecutor
from typing import Optional, Tuple, List
from PIL import Image
import torch
import numpy as np
from ..utils.image_utils import tensor_to_pil, pil_to_tensor
from ..utils.file_utils import (
ImageInfo,
load_images_from_folder,
pair_images_indexed,
pair_images_cartesian,
generate_output_filename,
save_image
)
from ..clients.gemini_client import GeminiAPIClient
from ..models_config import get_enabled_models
# 导入 ComfyUI 原生进度条
try:
from comfy.utils import ProgressBar
PROGRESS_BAR_AVAILABLE = True
except ImportError:
PROGRESS_BAR_AVAILABLE = False
print("⚠️ BatchNanoBananaPro: comfy.utils.ProgressBar 不可用,将只使用终端进度显示")
class BatchNanoBananaPro:
"""
批量 Nano Banana Pro 节点
功能:
- 从多个文件夹加载图片
- 支持三种配对模式:
* 1:1 - 索引配对(文件夹之间按位置配对)
* 1*N - 笛卡尔积配对(所有可能组合)
* 不配对 - 固定参考图模式(文件夹图片依次与所有参考图组合)
- 批量调用 API 生成图像
- 智能命名保存(保留原始文件名)
- 并发控制(默认最大 100)
注意:
- 「不配对」模式只支持单个文件夹
- 支持的模型列表从 models_config.py 动态加载
- 要添加/禁用模型,请编辑 models_config.py 文件
"""
# 支持的模型列表(从配置文件动态加载)
MODELS = None # 将在 INPUT_TYPES 中动态获取
# 支持的宽高比列表
ASPECT_RATIOS = [
"1:1", "4:3", "3:4", "16:9", "9:16",
"2:3", "3:2", "4:5", "5:4", "21:9"
]
# 支持的分辨率列表
RESOLUTIONS = ["1K", "2K", "4K"]
# 配对模式
PAIRING_MODES = ["1:1", "1*N", "不配对"]
def __init__(self):
"""初始化节点"""
self.client = None
def resize_to_megapixels(
self,
image: Image.Image,
target_megapixels: float
) -> Image.Image:
"""
将图像缩放到指定的总像素数,保持纵横比
Args:
image: PIL Image 对象
target_megapixels: 目标像素数(百万像素)
Returns:
缩放后的 PIL Image
Example:
>>> resized = self.resize_to_megapixels(img, 2.0) # 缩放到2百万像素
"""
# 计算当前像素数
current_pixels = image.width * image.height
target_pixels = int(target_megapixels * 1_000_000)
# 如果当前像素数已经接近目标,则不缩放
if abs(current_pixels - target_pixels) / target_pixels < 0.05:
return image
# 计算缩放比例
scale = (target_pixels / current_pixels) ** 0.5
# 计算新尺寸
new_width = int(image.width * scale)
new_height = int(image.height * scale)
# 确保至少为1像素
new_width = max(1, new_width)
new_height = max(1, new_height)
# 使用 Lanczos 重采样
resized_image = image.resize((new_width, new_height), Image.Resampling.LANCZOS)
return resized_image
@classmethod
def INPUT_TYPES(cls):
"""
定义输入参数
ComfyUI 节点规范:
- required: 必选参数
- optional: 可选参数
"""
# 从配置文件动态获取启用的模型列表
enabled_models = get_enabled_models()
# 如果没有启用的模型,使用空列表(会导致节点不可用,提示用户配置)
if not enabled_models:
enabled_models = ["请在 models_config.py 中启用至少一个模型"]
# 创建9个独立的图像输入
optional_inputs = {}
for i in range(1, 10): # 1-9
optional_inputs[f"参考图{i}"] = ("IMAGE",)
return {
"required": {
"prompt": ("STRING", {
"default": "一个中国女子的OOTD",
"multiline": True
}),
"模型": (enabled_models, {
"default": enabled_models[0]
}),
"宽高比": (cls.ASPECT_RATIOS, {
"default": "1:1"
}),
"分辨率": (cls.RESOLUTIONS, {
"default": "2K"
}),
"像素缩放": ("BOOLEAN", {
"default": False
}),
"分辨率像素": ("FLOAT", {
"default": 1.0,
"min": 0.1,
"max": 100.0,
"step": 0.1,
"display": "number"
}),
"seed": ("INT", {
"default": 0,
"min": 0,
"max": 0xffffffffffffffff
}),
"文件夹1": ("STRING", {
"default": "",
"multiline": False
}),
"文件夹2": ("STRING", {
"default": "",
"multiline": False
}),
"文件夹3": ("STRING", {
"default": "",
"multiline": False
}),
"文件夹4": ("STRING", {
"default": "",
"multiline": False
}),
"保存路径": ("STRING", {
"default": "",
"multiline": False
}),
"图片配对模式": (cls.PAIRING_MODES, {
"default": "不配对"
})
},
"optional": optional_inputs
}
# 返回值类型
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("输出图像",)
# 执行函数名
FUNCTION = "process_batch"
# 节点分类
CATEGORY = "image/batch"
def _load_folders(
self,
folder1: str,
folder2: Optional[str],
folder3: Optional[str],
folder4: Optional[str],
enable_scaling: bool,
target_megapixels: float
) -> List[List[ImageInfo]]:
"""
加载所有文件夹中的图片
Args:
folder1-4: 文件夹路径
enable_scaling: 是否启用像素缩放
target_megapixels: 目标像素数(百万像素)
Returns:
图片列表的列表
"""
folders = [folder1, folder2, folder3, folder4]
all_images = []
for i, folder in enumerate(folders, 1):
if folder and folder.strip():
try:
images = load_images_from_folder(folder)
if images:
# 应用像素缩放
if enable_scaling:
scaled_images = []
for img_info in images:
scaled_img = self.resize_to_megapixels(
img_info.image,
target_megapixels
)
# 创建新的 ImageInfo,保留其他元数据
scaled_info = ImageInfo(
image=scaled_img,
filename=img_info.filename,
extension=img_info.extension,
source_path=img_info.source_path
)
scaled_images.append(scaled_info)
images = scaled_images
all_images.append(images)
print(f"BatchNanoBananaPro: 文件夹{i} 加载了 {len(images)} 张图片")
else:
print(f"BatchNanoBananaPro: 文件夹{i} 为空或没有有效图片")
except ValueError as e:
print(f"BatchNanoBananaPro: 文件夹{i} 加载失败 - {e}")
return all_images
def _create_pairs(
self,
image_lists: List[List[ImageInfo]],
pairing_mode: str,
manual_images: Optional[List[ImageInfo]] = None
) -> List[Tuple[ImageInfo, ...]]:
"""
根据配对模式创建图片组合
Args:
image_lists: 从文件夹加载的图片列表
pairing_mode: 配对模式 (1:1, 1*N, 不配对)
manual_images: 手动输入的参考图
Returns:
配对后的元组列表
Raises:
ValueError: 不配对模式下填入多个文件夹时
"""
# === 新模式:不配对 ===
if pairing_mode == "不配对":
# 验证:只支持单个文件夹
if len(image_lists) > 1:
raise ValueError("「不配对」模式只支持单个文件夹,请清空其他文件夹路径")
# 场景1:有文件夹 + 有参考图
if image_lists and manual_images:
folder_images = image_lists[0]
# 每张文件夹图片 + 所有参考图
pairs = []
for img in folder_images:
pair = (img,) + tuple(manual_images)
pairs.append(pair)
return pairs
# 场景2:有文件夹 + 无参考图
elif image_lists:
# 每张图片单独成组
return [(img,) for img in image_lists[0]]
# 场景3:无文件夹 + 有参考图
elif manual_images:
# 每张参考图单独成组
return [(img,) for img in manual_images]
else:
return []
# === 原有逻辑:1:1 和 1*N ===
# 如果有手动参考图,添加到列表中(所有参考图作为一个列表)
if manual_images:
image_lists.append(manual_images)
if not image_lists:
return []
# 如果只有一个列表,直接返回每个图片作为单元素元组
if len(image_lists) == 1:
return [(img,) for img in image_lists[0]]
# 根据配对模式选择配对函数
if pairing_mode == "1:1":
pairs = pair_images_indexed(*image_lists)
else: # 1*N
pairs = pair_images_cartesian(*image_lists)
return pairs
async def _generate_single_task(
self,
client: GeminiAPIClient,
session: aiohttp.ClientSession,
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
images: List[ImageInfo],
output_folder: str,
task_index: int
) -> dict:
"""
执行单个生成任务
Args:
client: API 客户端
session: aiohttp 会话
prompt: 提示词
model: 模型名称
resolution: 分辨率
aspect_ratio: 宽高比
images: 输入图片列表
output_folder: 输出文件夹
task_index: 任务索引
Returns:
包含结果信息的字典
"""
result = {
"task_index": task_index,
"success": False,
"generated_count": 0,
"saved_files": [],
"error": None
}
try:
# 准备输入图片
input_pil_images = [info.image for info in images]
# 调用 API 生成图片(固定生成1次)
generated_images = []
try:
gen_result = await client.generate_single_async(
prompt=prompt,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
images=input_pil_images,
session=session
)
if gen_result:
generated_images.extend(gen_result)
except Exception as e:
error_msg = str(e)
print(f"BatchNanoBananaPro: 任务 {task_index + 1} 生成失败 - {error_msg}")
result["error"] = error_msg
# 保存生成的图片
for i, gen_img in enumerate(generated_images):
# 使用任务索引作为唯一标识,确保并发安全
output_path = generate_output_filename(
source_images=list(images),
batch_index=i,
output_folder=output_folder,
extension=".png",
task_id=f"task{task_index}"
)
save_image(gen_img, output_path)
result["saved_files"].append(output_path)
# 只有生成了图片才标记为成功
if len(generated_images) > 0:
result["success"] = True
result["generated_count"] = len(generated_images)
except Exception as e:
result["error"] = str(e)
return result
async def _process_batch_async(
self,
pairs: List[Tuple[ImageInfo, ...]],
prompt: str,
model: str,
resolution: str,
aspect_ratio: str,
output_folder: str,
pbar=None
) -> List[dict]:
"""
异步批量处理所有任务
Args:
pairs: 配对后的图片组合
prompt: 提示词
model: 模型名称
resolution: 分辨率
aspect_ratio: 宽高比
output_folder: 输出文件夹
Returns:
所有任务的结果列表
"""
if self.client is None:
self.client = GeminiAPIClient()
# 固定最大并发数为 100
max_concurrent = 100
total_tasks = len(pairs)
all_results = []
completed = 0
success_count = 0
fail_count = 0
# 计算分批数量
num_batches = math.ceil(total_tasks / max_concurrent)
# 进度打印配置:任务数 >= 50 时,额外显示百分比里程碑
show_milestone = total_tasks >= 50
milestones = [0.2, 0.4, 0.6, 0.8, 1.0] # 20%, 40%, 60%, 80%, 100%
milestone_index = 0
if num_batches > 1:
print(f"BatchNanoBananaPro: 任务数 {total_tasks} 超过并发上限 {max_concurrent},将分 {num_batches} 批执行")
connector = aiohttp.TCPConnector(limit=0, limit_per_host=0)
async with aiohttp.ClientSession(connector=connector) as session:
for batch_idx in range(num_batches):
start_idx = batch_idx * max_concurrent
end_idx = min(start_idx + max_concurrent, total_tasks)
batch_pairs = pairs[start_idx:end_idx]
if num_batches > 1:
print(f"BatchNanoBananaPro: 执行第 {batch_idx + 1}/{num_batches} 批 ({start_idx + 1}-{end_idx})...")
# 创建当前批次的任务
tasks = []
for i, pair in enumerate(batch_pairs):
task = asyncio.create_task(
self._generate_single_task(
client=self.client,
session=session,
prompt=prompt,
model=model,
resolution=resolution,
aspect_ratio=aspect_ratio,
images=list(pair),
output_folder=output_folder,
task_index=start_idx + i
)
)
tasks.append(task)
# 使用 as_completed 实时获取完成的任务
for coro in asyncio.as_completed(tasks):
result_data = None
try:
result = await coro
if isinstance(result, Exception):
result_data = {
"success": False,
"error": str(result),
"generated_count": 0,
"saved_files": []
}
all_results.append(result_data)
else:
result_data = result
all_results.append(result)
except Exception as e:
result_data = {
"success": False,
"error": str(e),
"generated_count": 0,
"saved_files": []
}
all_results.append(result_data)
completed += 1
# 根据成功/失败状态打印不同信息
if result_data and result_data.get("success", False):
success_count += 1
print(f"BatchNanoBananaPro: 任务 {completed}/{total_tasks} 成功 ✓")
else:
fail_count += 1
# 提取错误信息的第一行
error_msg = result_data.get("error", "未知错误") if result_data else "未知错误"
# 截取第一行或前50个字符
if '\n' in error_msg:
error_msg = error_msg.split('\n')[0]
if len(error_msg) > 50:
error_msg = error_msg[:50] + "..."
print(f"BatchNanoBananaPro: 任务 {completed}/{total_tasks} 失败 ✗ - {error_msg}")
# 更新 ComfyUI 原生进度条
if pbar is not None:
pbar.update(1)
# 大任务额外显示百分比里程碑
if show_milestone and milestone_index < len(milestones):
progress = completed / total_tasks
if progress >= milestones[milestone_index]:
percentage = int(milestones[milestone_index] * 100)
print(f"BatchNanoBananaPro: >>> 进度 {percentage}% <<<")
milestone_index += 1
return all_results
def process_batch(
self,
prompt: str,
文件夹1: str,
文件夹2: str,
文件夹3: str,
文件夹4: str,
像素缩放: bool,
分辨率像素: float,
seed: int,
保存路径: str,
图片配对模式: str,
模型: str,
宽高比: str,
分辨率: str,
**kwargs
) -> Tuple[torch.Tensor]:
"""
批量处理图像生成任务
Args:
prompt: 提示词
文件夹1-4: 图片文件夹路径
像素缩放: 是否启用像素缩放
分辨率像素: 目标像素数(百万像素)
seed: 随机种子
保存路径: 输出保存路径
图片配对模式: 1:1 或 1*N
模型: 模型名称
宽高比: 输出宽高比
分辨率: 输出分辨率
**kwargs: 动态参考图输入 (参考图1-9)
Returns:
输出图像张量
"""
start_time = time.time()
try:
# 设置随机种子(用于本地随机操作)
random.seed(seed)
np.random.seed(seed % (2**32))
# 验证保存路径
if not 保存路径 or not 保存路径.strip():
raise ValueError("请提供保存路径")
# 加载文件夹图片
print("BatchNanoBananaPro: 开始加载图片...")
image_lists = self._load_folders(
文件夹1, 文件夹2, 文件夹3, 文件夹4,
像素缩放, 分辨率像素
)
# 处理独立的参考图输入
manual_images = []
for i in range(1, 10): # 1-9
key = f"参考图{i}"
if key in kwargs and kwargs[key] is not None:
pil_images = tensor_to_pil(kwargs[key])
for j, img in enumerate(pil_images):
# 如果启用像素缩放,也对参考图进行缩放
if 像素缩放:
img = self.resize_to_megapixels(img, 分辨率像素)
manual_images.append(
ImageInfo(
image=img,
filename=f"manual_{i}_{j}",
extension=".png",
source_path=""
)
)
if manual_images:
print(f"BatchNanoBananaPro: 加载了 {len(manual_images)} 张参考图")
# 验证是否有图片
total_folder_images = sum(len(lst) for lst in image_lists)
total_manual_images = len(manual_images)
if total_folder_images == 0 and total_manual_images == 0:
raise ValueError("未找到任何图片,请检查文件夹路径或提供参考图")
# 创建配对
print(f"BatchNanoBananaPro: 使用 {图片配对模式} 模式创建配对...")
pairs = self._create_pairs(image_lists, 图片配对模式, manual_images if manual_images else None)
if not pairs:
raise ValueError("配对结果为空,请检查输入")
total_tasks = len(pairs)
print(f"BatchNanoBananaPro: 共 {total_tasks} 组配对")
# 创建 ComfyUI 原生进度条
pbar = None
if PROGRESS_BAR_AVAILABLE:
pbar = ProgressBar(total_tasks)
# 初始化 API 客户端
if self.client is None:
try:
self.client = GeminiAPIClient()
except ValueError as e:
raise ValueError(f"初始化 API 客户端失败: {str(e)}")
# 执行批量生成
print("BatchNanoBananaPro: 开始批量生成...")
# 在新线程中运行异步代码,避免事件循环冲突
def run_async_in_thread():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
return loop.run_until_complete(
self._process_batch_async(
pairs=pairs,
prompt=prompt,
model=模型,
resolution=分辨率,
aspect_ratio=宽高比,
output_folder=保存路径,
pbar=pbar
)
)
finally:
loop.close()
# 使用线程池在新线程中运行事件循环
with ThreadPoolExecutor(max_workers=1) as executor:
future = executor.submit(run_async_in_thread)
results = future.result()
# 统计结果
success_count = sum(1 for r in results if r.get("success", False))
fail_count = len(results) - success_count
total_generated = sum(r.get("generated_count", 0) for r in results)
all_saved_files = []
for r in results:
all_saved_files.extend(r.get("saved_files", []))
elapsed = time.time() - start_time
# 精简统计信息
print("=" * 50)
print(f"BatchNanoBananaPro 处理完成 | 总耗时: {elapsed:.2f}s | 成功: {success_count}/{total_tasks} | 生成: {total_generated}")
print(f"保存路径: {保存路径}")
# 失败详情(如果有)
failed_results = [r for r in results if not r.get("success", False)]
if failed_results:
# 收集失败任务的索引
failed_indices = [str(r.get('task_index', '?') + 1) for r in failed_results[:5]]
failed_str = ",".join(failed_indices)
if len(failed_results) > 5:
failed_str += f"... (共{len(failed_results)}个)"
# 显示第一个失败原因作为示例
first_error = failed_results[0].get('error', '未知错误')
print(f"失败 {len(failed_results)}个: 任务{failed_str} - {first_error}")
# 收集所有生成的图片
output_images = []
for file_path in all_saved_files:
try:
img = Image.open(file_path)
output_images.append(img)
except Exception as e:
print(f"BatchNanoBananaPro: 无法加载图片 {file_path} - {e}")
# 如果没有生成成功的图片,创建一个占位图
if not output_images:
placeholder = Image.new('RGB', (512, 512), color=(128, 128, 128))
output_images = [placeholder]
# 转换为张量
output_tensor = pil_to_tensor(output_images)
return (output_tensor,)
except ValueError as e:
# 检测是否为授权错误
if str(e) == "未授权!":
print("请联系作者授权后方可使用!")
else:
print(f"BatchNanoBananaPro: 输入错误 - {str(e)}")
raise
except RuntimeError as e:
print(f"BatchNanoBananaPro: 运行时错误 - {str(e)}")
raise
except Exception as e:
print(f"BatchNanoBananaPro: 未知错误 - {str(e)}")
raise
finally:
# 无论成功或失败,都尝试查询余额
if self.client is not None:
try:
balance_data = self.client.query_balance_sync()
balance_info = self.client.format_balance_info(balance_data)
print(f"{balance_info}")
print("=" * 50)
except Exception as e:
print(f"⚠️ 余额查询失败 - {str(e)}")
print("=" * 50)
+332
View File
@@ -0,0 +1,332 @@
"""
Google Gemini 节点
ComfyUI 自定义节点,用于调用 Gemini Flash 模型进行多模态文本生成
"""
import base64
import os
import time
from typing import Dict, List, Optional, Tuple
import torch
from ..utils.image_utils import tensor_to_pil, encode_image_to_base64
from ..clients.gemini_flash_client import GeminiFlashClient
from ..models_config import get_enabled_flash_models
# 支持的视频 MIME 类型映射
VIDEO_MIME_TYPES = {
".mp4": "video/mp4",
".mpeg": "video/mpeg",
".mpg": "video/mpg",
".mov": "video/mov",
".avi": "video/avi",
".flv": "video/x-flv",
".webm": "video/webm",
".wmv": "video/wmv",
".3gp": "video/3gpp",
".3gpp": "video/3gpp"
}
class GoogleGemini:
"""
Google Gemini 节点
功能:
- 支持多个 Gemini Flash 模型
- 支持图片和视频输入
- 支持系统指令
- 支持不同思考深度(不思考/高)
- 输出生成的文本内容
"""
# 支持的思考深度选项
THINKING_DEPTHS = ["不思考", ""]
def __init__(self):
"""初始化节点"""
self.client = None
@classmethod
def INPUT_TYPES(cls):
"""
定义输入参数
"""
# 从配置获取启用的模型列表
enabled_models = get_enabled_flash_models()
default_model = enabled_models[0] if enabled_models else "gemini-3-flash-preview"
return {
"required": {
"模型": (enabled_models, {
"default": default_model
}),
"提示词": ("STRING", {
"default": "",
"multiline": True
}),
"思考深度": (cls.THINKING_DEPTHS, {
"default": "不思考"
})
},
"optional": {
"系统指令": ("STRING", {
"default": "",
"multiline": True
}),
"图片": ("IMAGE",),
"视频": ("VIDEO",)
}
}
# 返回值类型
RETURN_TYPES = ("STRING", "STRING")
RETURN_NAMES = ("主要内容", "思考内容")
# 执行函数名
FUNCTION = "generate"
# 节点分类
CATEGORY = "text/generation"
# 允许输出到 UI
OUTPUT_NODE = True
def _prepare_image_data(
self,
images: Optional[torch.Tensor]
) -> Optional[List[Dict[str, str]]]:
"""
准备图片数据
Args:
images: ComfyUI 图片张量 [B, H, W, C]
Returns:
图片数据列表,每个元素包含 mime_type 和 data
"""
if images is None:
return None
image_data = []
pil_images = tensor_to_pil(images)
for img in pil_images:
b64_str = encode_image_to_base64(img)
image_data.append({
"mime_type": "image/png",
"data": b64_str
})
return image_data if image_data else None
def _prepare_video_data(
self,
video
) -> Optional[Dict[str, str]]:
"""
准备视频数据
ComfyUI VIDEO 类型包含视频文件路径信息。
读取视频文件并转换为 base64。
Args:
video: ComfyUI VIDEO 类型数据
Returns:
视频数据字典,包含 mime_type 和 data
"""
if video is None:
return None
# VIDEO 类型通常是一个字典,包含 'video' 键指向文件路径
# 或者直接是文件路径字符串
video_path = None
if isinstance(video, dict):
# 尝试获取视频路径
video_path = video.get("video") or video.get("path") or video.get("file")
elif isinstance(video, str):
video_path = video
elif hasattr(video, "video"):
video_path = video.video
if not video_path or not os.path.exists(video_path):
print(f"Google Gemini: 视频文件不存在或路径无效: {video_path}")
return None
# 获取文件扩展名和 MIME 类型
_, ext = os.path.splitext(video_path)
ext = ext.lower()
mime_type = VIDEO_MIME_TYPES.get(ext, "video/mp4")
# 检查文件大小(限制 20MB
file_size = os.path.getsize(video_path)
if file_size > 20 * 1024 * 1024:
raise ValueError(
f"视频文件过大 ({file_size / 1024 / 1024:.2f}MB)"
f"请使用不超过 20MB 的视频文件"
)
# 读取并编码视频
try:
with open(video_path, "rb") as f:
video_bytes = f.read()
b64_str = base64.b64encode(video_bytes).decode("utf-8")
return {
"mime_type": mime_type,
"data": b64_str
}
except Exception as e:
print(f"Google Gemini: 读取视频文件失败 - {str(e)}")
return None
def _parse_dual_output(self, raw_response: Dict) -> Tuple[str, str]:
"""
解析包含思考内容和主要内容的响应
Args:
raw_response: API 原始响应字典
Returns:
(主要内容, 思考内容)
"""
candidates = raw_response.get("candidates", [])
if not candidates:
return ("", "")
parts = candidates[0].get("content", {}).get("parts", [])
thought_text = ""
main_text = ""
for part in parts:
if part.get("thought") is True:
# 思考部分
thought_text = part.get("text", "")
elif "thoughtSignature" in part or "text" in part:
# 主要内容
main_text = part.get("text", "")
return (main_text, thought_text)
def generate(
self,
模型: str,
提示词: str,
思考深度: str,
系统指令: Optional[str] = None,
图片: Optional[torch.Tensor] = None,
视频=None
) -> Tuple[str]:
"""
生成文本
Args:
模型: 使用的模型名称
提示词: 用户提示词
思考深度: 思考深度选项
系统指令: 系统级指令
图片: 输入图片
视频: 输入视频
Returns:
生成的文本 (STRING,)
"""
start_time = time.time()
try:
# 初始化 API 客户端
if self.client is None:
try:
self.client = GeminiFlashClient()
except ValueError as e:
raise ValueError(f"初始化失败: {str(e)}")
# 准备图片数据
image_data = self._prepare_image_data(图片)
if image_data:
print(f"Google Gemini: 输入 {len(image_data)} 张图片")
# 准备视频数据
video_data = self._prepare_video_data(视频)
if video_data:
print(f"Google Gemini: 输入视频 ({video_data['mime_type']})")
# 构建输入描述
input_desc = []
if 提示词:
input_desc.append("文本")
if image_data:
input_desc.append(f"{len(image_data)}张图片")
if video_data:
input_desc.append("视频")
print(f"Google Gemini: 模型 = {模型}")
print(f"Google Gemini: 多模态输入 ({', '.join(input_desc)})")
print(f"Google Gemini: 思考深度 = {思考深度}")
print(f"Google Gemini: 发送请求...")
# 获取端点和构建请求体
endpoint = self.client.get_endpoint(model=模型, thinking_depth=思考深度)
request_body = self.client.build_request_body(
prompt=提示词,
system_instruction=系统指令,
image_data=image_data,
video_data=video_data
)
# 根据是否有视频设置超时
timeout = 300 if video_data else 180
# 调用底层 API 获取原始响应
async def get_raw_response():
return await self.client.request_async(
endpoint,
request_body,
session=None,
timeout=timeout
)
# 在独立线程中执行异步请求
raw_response = self.client.run_async_in_thread(get_raw_response())
# 计算耗时
elapsed = time.time() - start_time
# 解析响应,分离主要内容和思考内容
main_text, thought_text = self._parse_dual_output(raw_response)
# 输出信息
print(f"Google Gemini: 生成完成 (耗时: {elapsed:.2f}s)")
print(f"Google Gemini: 主要内容长度: {len(main_text)} 字符")
print(f"Google Gemini: 思考内容长度: {len(thought_text)} 字符")
# 输出预览
if main_text:
preview = main_text[:100] + "..." if len(main_text) > 100 else main_text
print(f"Google Gemini: 主要内容预览: {preview}")
return (main_text, thought_text)
except ValueError as e:
# 检测是否为授权错误
if str(e) == "未授权!":
print("请联系作者授权后方可使用!")
else:
print(f"Google Gemini: 输入错误 - {str(e)}")
raise
except RuntimeError as e:
print(f"Google Gemini: API 错误 - {str(e)}")
raise
except Exception as e:
print(f"Google Gemini: 未知错误 - {str(e)}")
raise
+381
View File
@@ -0,0 +1,381 @@
"""
Nano Banana Pro 节点
ComfyUI 自定义节点,用于调用 Gemini 3 Pro 模型生成图像
"""
import time
import random
from typing import Optional, Tuple
import torch
import numpy as np
from PIL import Image
from ..utils.image_utils import tensor_to_pil, pil_to_tensor, parse_batch_prompts
from ..clients.gemini_client import GeminiAPIClient
from ..models_config import get_enabled_models, get_model_description
# 导入 ComfyUI 原生进度条
try:
from comfy.utils import ProgressBar
PROGRESS_BAR_AVAILABLE = True
except ImportError:
PROGRESS_BAR_AVAILABLE = False
print("⚠️ NanoBananaPro: comfy.utils.ProgressBar 不可用,将只使用终端进度显示")
class NanoBananaPro:
"""
Nano Banana Pro 节点
功能:
- 文生图:基于提示词生成图像
- 图生图:基于输入图像和提示词生成新图像
- 批量生成:支持并发生成多张图像
注意:
- 支持的模型列表从 models_config.py 动态加载
- 要添加/禁用模型,请编辑 models_config.py 文件
"""
# 支持的模型列表(从配置文件动态加载)
MODELS = None # 将在 INPUT_TYPES 中动态获取
# 支持的宽高比列表
ASPECT_RATIOS = [
"1:1", "4:3", "3:4", "16:9", "9:16",
"2:3", "3:2", "4:5", "5:4", "21:9"
]
# 支持的分辨率列表
RESOLUTIONS = ["1K", "2K", "4K"]
def __init__(self):
"""初始化节点"""
self.client = None
@classmethod
def INPUT_TYPES(cls):
"""
定义输入参数
ComfyUI 节点规范:
- required: 必选参数
- optional: 可选参数
"""
# 从配置文件动态获取启用的模型列表
enabled_models = get_enabled_models()
# 如果没有启用的模型,使用空列表(会导致节点不可用,提示用户配置)
if not enabled_models:
enabled_models = ["请在 models_config.py 中启用至少一个模型"]
# 创建9个独立的图像输入
optional_inputs = {}
for i in range(1, 10): # 1-9
optional_inputs[f"参考图{i}"] = ("IMAGE",)
return {
"required": {
"prompt": ("STRING", {
"default": "一个中国女子的OOTD",
"multiline": True
}),
"模型": (enabled_models, {
"default": enabled_models[0]
}),
"宽高比": (cls.ASPECT_RATIOS, {
"default": "1:1"
}),
"分辨率": (cls.RESOLUTIONS, {
"default": "2K"
}),
"生图数量": ("INT", {
"default": 1,
"min": 1,
"max": 1000,
"step": 1
}),
"像素缩放": ("BOOLEAN", {
"default": False
}),
"分辨率像素": ("FLOAT", {
"default": 1.0,
"min": 0.1,
"max": 100.0,
"step": 0.1,
"display": "number"
}),
"seed": ("INT", {
"default": 0,
"min": 0,
"max": 0xffffffffffffffff
})
},
"optional": optional_inputs
}
# 返回值类型
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("输出图像",)
# 执行函数名
FUNCTION = "generate"
# 节点分类
CATEGORY = "image/generation"
def resize_to_megapixels(
self,
image: Image.Image,
target_megapixels: float
) -> Image.Image:
"""
将图像缩放到指定的总像素数,保持纵横比
Args:
image: PIL Image 对象
target_megapixels: 目标像素数(百万像素)
Returns:
缩放后的 PIL Image
Example:
>>> resized = self.resize_to_megapixels(img, 2.0) # 缩放到2百万像素
"""
# 计算当前像素数
current_pixels = image.width * image.height
target_pixels = int(target_megapixels * 1_000_000)
# 如果当前像素数已经接近目标,则不缩放
if abs(current_pixels - target_pixels) / target_pixels < 0.05:
return image
# 计算缩放比例
scale = (target_pixels / current_pixels) ** 0.5
# 计算新尺寸
new_width = int(image.width * scale)
new_height = int(image.height * scale)
# 确保至少为1像素
new_width = max(1, new_width)
new_height = max(1, new_height)
# 使用 Lanczos 重采样
resized_image = image.resize((new_width, new_height), Image.Resampling.LANCZOS)
return resized_image
def validate_inputs(
self,
images: Optional[torch.Tensor],
batch_size: int
) -> None:
"""
验证输入参数
Args:
images: 输入图像张量(可选)
batch_size: 批次大小
Raises:
ValueError: 如果输入参数不合法
"""
# 检查图像数量
if images is not None:
num_images = images.shape[0]
if num_images > 14:
raise ValueError(
f"输入图像数量 {num_images} 超过限制 14 张,请减少输入图像数量"
)
# 检查批次大小
if batch_size < 1 or batch_size > 1000:
raise ValueError(
f"批次大小 {batch_size} 超出范围 [1, 1000]"
)
def generate(
self,
prompt: str,
模型: str,
宽高比: str,
分辨率: str,
生图数量: int,
像素缩放: bool,
分辨率像素: float,
seed: int,
**kwargs
) -> Tuple[torch.Tensor]:
"""
生成图像
Args:
prompt: 提示词
模型: 模型名称
宽高比: 宽高比
分辨率: 分辨率
生图数量: 批次大小
像素缩放: 是否启用像素缩放
分辨率像素: 目标像素数(百万像素)
seed: 随机种子
**kwargs: 动态参考图输入 (参考图1-9)
Returns:
生成的图像张量 (IMAGE,)
"""
start_time = time.time()
# 创建 ComfyUI 原生进度条
pbar = None
if PROGRESS_BAR_AVAILABLE:
pbar = ProgressBar(生图数量)
try:
# 设置随机种子(用于本地随机操作)
random.seed(seed)
np.random.seed(seed % (2**32))
# 初始化 API 客户端
if self.client is None:
try:
self.client = GeminiAPIClient()
except ValueError as e:
raise ValueError(f"初始化失败: {str(e)}")
# 收集独立输入的参考图
input_images = []
for i in range(1, 10): # 1-9
key = f"参考图{i}"
if key in kwargs and kwargs[key] is not None:
pil_imgs = tensor_to_pil(kwargs[key])
input_images.extend(pil_imgs)
# 验证输入图像数量
if input_images:
if len(input_images) > 14:
raise ValueError(
f"输入图像数量 {len(input_images)} 超过限制 14 张,请减少输入图像数量"
)
# 应用像素缩放(如果启用)
if input_images and 像素缩放:
scaled_images = []
for img in input_images:
scaled = self.resize_to_megapixels(img, 分辨率像素)
scaled_images.append(scaled)
input_images = scaled_images
print(f"Nano Banana Pro: 已缩放 {len(scaled_images)} 张图像到 {分辨率像素}M 像素")
# 转换为 API 所需的格式
if input_images:
print(f"Nano Banana Pro: 图生图模式 (输入 {len(input_images)} 张图像)")
# 解析批量提示词
batch_prompts = parse_batch_prompts(prompt)
# 统计变量
success_count = 0
fail_count = 0
# 进度回调 - 实时显示每个任务的完成状态,并更新 ComfyUI 进度条
def progress_callback(current, total, success, error_msg=None):
nonlocal success_count, fail_count
if success:
success_count += 1
print(f"Nano Banana Pro: ✓ [{current}/{total}] 第 {success_count} 张生成成功")
else:
fail_count += 1
error_brief = error_msg[:50] + "..." if error_msg and len(error_msg) > 50 else error_msg
print(f"Nano Banana Pro: ✗ [{current}/{total}] 生成失败 - {error_brief}")
# 更新 ComfyUI 原生进度条
if pbar is not None:
pbar.update(1)
# 根据是否有批量提示词选择生成模式
if batch_prompts:
# 批量提示词模式
num_prompts = len(batch_prompts)
total_images = num_prompts * 生图数量
print(f"Nano Banana Pro: 批量提示词模式 ({num_prompts} 个提示词 × {生图数量} 张/提示词 = {total_images} 张图)")
print(f"Nano Banana Pro: 发送请求")
print(f"Nano Banana Pro: 生图中...")
# 重新创建进度条以匹配实际总数
if pbar is not None:
pbar = ProgressBar(total_images)
generated_images = self.client.generate_multi_prompts_sync(
prompts=batch_prompts,
model=模型,
resolution=分辨率,
aspect_ratio=宽高比,
images_per_prompt=生图数量,
images=input_images,
progress_callback=progress_callback
)
if fail_count > 0:
print(f"Nano Banana Pro: 生图完成 (成功: {success_count}, 失败: {fail_count})")
else:
print(f"Nano Banana Pro: 全部生图成功!")
else:
# 单提示词模式
print(f"Nano Banana Pro: {'图生图' if input_images else '文生图'}模式")
print(f"Nano Banana Pro: 发送请求")
print(f"Nano Banana Pro: 生图中...")
generated_images = self.client.generate_sync(
prompt=prompt,
model=模型,
resolution=分辨率,
aspect_ratio=宽高比,
batch_size=生图数量,
images=input_images,
progress_callback=progress_callback
)
if fail_count > 0:
print(f"Nano Banana Pro: 生图完成 (成功: {success_count}, 失败: {fail_count})")
else:
print(f"Nano Banana Pro: 全部生图成功!")
# 转换输出图像
output_tensor = pil_to_tensor(generated_images)
# 计算耗时
elapsed = time.time() - start_time
print(f"Nano Banana Pro: 完成生图 (耗时: {elapsed:.2f}s, 成功生成 {len(generated_images)} 张图像)")
return (output_tensor,)
except ValueError as e:
# 检测是否为授权错误
if str(e) == "未授权!":
print("请联系作者授权后方可使用!")
else:
# 用户输入错误
print(f"Nano Banana Pro: 输入错误 - {str(e)}")
raise
except RuntimeError as e:
# API 或网络错误
print(f"Nano Banana Pro: API 错误 - {str(e)}")
raise
except Exception as e:
# 其他未知错误
print(f"Nano Banana Pro: 未知错误 - {str(e)}")
raise
finally:
# 无论成功或失败,都尝试查询余额
if self.client is not None:
try:
balance_data = self.client.query_balance_sync()
balance_info = self.client.format_balance_info(balance_data)
print(f"Nano Banana Pro: {balance_info}")
except Exception as e:
print(f"Nano Banana Pro: ⚠️ 余额查询失败 - {str(e)}")
+3
View File
@@ -0,0 +1,3 @@
aiohttp>=3.9.0
Pillow>=10.0.0
requests>=2.31.0
+97
View File
@@ -0,0 +1,97 @@
@echo off
chcp 65001 > nul
echo ====================================
echo Comfyui_o1key 插件更新工具
echo ====================================
echo.
:: 检查是否在 Git 仓库中
if not exist ".git" (
echo [错误] 当前目录不是 Git 仓库
echo 请确保插件是通过 git clone 安装的
pause
exit /b 1
)
:: 保存当前版本
if exist "version.txt" (
set /p OLD_VERSION=<version.txt
echo 当前版本: %OLD_VERSION%
) else (
set OLD_VERSION=未知
echo 当前版本: 未知
)
echo.
echo [1/4] 检查远程更新...
git fetch origin
:: 检查是否有更新
git status -uno | findstr "Your branch is behind" > nul
if %errorlevel% equ 0 (
echo 发现新版本!
) else (
echo 已是最新版本
echo.
choice /C YN /M "是否继续检查依赖更新?"
if errorlevel 2 goto :end
)
echo.
echo [2/4] 备份配置文件...
if exist ".config" (
copy /Y ".config" ".config.backup" > nul
echo 已备份 .config 到 .config.backup
)
echo.
echo [3/4] 拉取最新代码...
git pull origin main
if %errorlevel% neq 0 (
echo [错误] 代码更新失败,请检查网络连接或手动解决冲突
pause
exit /b 1
)
:: 恢复配置文件
if exist ".config.backup" (
copy /Y ".config.backup" ".config" > nul
del ".config.backup"
echo 已恢复配置文件
)
echo.
echo [4/4] 更新依赖包...
python -m pip install -r requirements.txt --upgrade --quiet
if %errorlevel% neq 0 (
echo [警告] 依赖包更新失败,请手动运行: pip install -r requirements.txt
)
echo.
echo ====================================
echo 更新完成!
echo ====================================
:: 显示新版本
if exist "version.txt" (
set /p NEW_VERSION=<version.txt
echo 新版本: %NEW_VERSION%
)
:: 显示最近更新日志
if exist "CHANGELOG.md" (
echo.
echo 最近更新内容:
echo -----------------------------------
powershell -Command "Get-Content CHANGELOG.md -TotalCount 20"
echo -----------------------------------
)
echo.
echo 请重启 ComfyUI 以使更改生效
echo.
pause
goto :end
:end
exit /b 0
+101
View File
@@ -0,0 +1,101 @@
#!/bin/bash
# 设置颜色输出
GREEN='\033[0;32m'
YELLOW='\033[1;33m'
RED='\033[0;31m'
NC='\033[0m' # No Color
echo "===================================="
echo "Comfyui_o1key 插件更新工具"
echo "===================================="
echo ""
# 检查是否在 Git 仓库中
if [ ! -d ".git" ]; then
echo -e "${RED}[错误] 当前目录不是 Git 仓库${NC}"
echo "请确保插件是通过 git clone 安装的"
exit 1
fi
# 保存当前版本
if [ -f "version.txt" ]; then
OLD_VERSION=$(cat version.txt)
echo "当前版本: $OLD_VERSION"
else
OLD_VERSION="未知"
echo "当前版本: 未知"
fi
echo ""
echo "[1/4] 检查远程更新..."
git fetch origin
# 检查是否有更新
LOCAL=$(git rev-parse @)
REMOTE=$(git rev-parse @{u})
if [ $LOCAL != $REMOTE ]; then
echo -e "${GREEN}发现新版本!${NC}"
else
echo -e "${GREEN}已是最新版本${NC}"
read -p "是否继续检查依赖更新?(y/n) " -n 1 -r
echo
if [[ ! $REPLY =~ ^[Yy]$ ]]; then
exit 0
fi
fi
echo ""
echo "[2/4] 备份配置文件..."
if [ -f ".config" ]; then
cp .config .config.backup
echo "已备份 .config 到 .config.backup"
fi
echo ""
echo "[3/4] 拉取最新代码..."
if git pull origin main; then
echo -e "${GREEN}代码更新成功${NC}"
else
echo -e "${RED}[错误] 代码更新失败,请检查网络连接或手动解决冲突${NC}"
exit 1
fi
# 恢复配置文件
if [ -f ".config.backup" ]; then
mv .config.backup .config
echo "已恢复配置文件"
fi
echo ""
echo "[4/4] 更新依赖包..."
if python3 -m pip install -r requirements.txt --upgrade --quiet; then
echo -e "${GREEN}依赖包更新成功${NC}"
else
echo -e "${YELLOW}[警告] 依赖包更新失败,请手动运行: pip install -r requirements.txt${NC}"
fi
echo ""
echo "===================================="
echo -e "${GREEN}更新完成!${NC}"
echo "===================================="
# 显示新版本
if [ -f "version.txt" ]; then
NEW_VERSION=$(cat version.txt)
echo "新版本: $NEW_VERSION"
fi
# 显示最近更新日志
if [ -f "CHANGELOG.md" ]; then
echo ""
echo "最近更新内容:"
echo "-----------------------------------"
head -n 20 CHANGELOG.md
echo "-----------------------------------"
fi
echo ""
echo "请重启 ComfyUI 以使更改生效"
echo ""
+39
View File
@@ -0,0 +1,39 @@
"""
工具模块
包含图像处理、配置管理、文件处理等通用工具函数
"""
from .image_utils import (
tensor_to_pil,
pil_to_tensor,
encode_image_to_base64,
decode_base64_to_pil
)
from .config import load_config, get_api_key
from .file_utils import (
ImageInfo,
load_images_from_folder,
pair_images_indexed,
pair_images_cartesian,
generate_output_filename,
generate_batch_output_filenames,
save_image,
get_folder_image_count
)
__all__ = [
'tensor_to_pil',
'pil_to_tensor',
'encode_image_to_base64',
'decode_base64_to_pil',
'load_config',
'get_api_key',
'ImageInfo',
'load_images_from_folder',
'pair_images_indexed',
'pair_images_cartesian',
'generate_output_filename',
'generate_batch_output_filenames',
'save_image',
'get_folder_image_count'
]
+114
View File
@@ -0,0 +1,114 @@
"""
配置管理模块
处理环境变量和 API 密钥管理
"""
import os
from typing import Dict, Optional
# 获取插件根目录
PLUGIN_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
CONFIG_FILE = os.path.join(PLUGIN_ROOT, ".config")
def load_config(config_path: Optional[str] = None) -> Dict[str, str]:
"""
从配置文件加载所有配置项
Args:
config_path: 配置文件路径,默认为插件目录下的 .config
Returns:
配置字典 {key: value}
Example:
>>> config = load_config()
>>> api_key = config.get('O1KEY_API_KEY')
"""
if config_path is None:
config_path = CONFIG_FILE
config = {}
if not os.path.exists(config_path):
return config
try:
with open(config_path, 'r', encoding='utf-8') as f:
for line in f:
line = line.strip()
# 跳过空行和注释
if not line or line.startswith('#'):
continue
# 解析 KEY=VALUE 格式
if '=' in line:
key, value = line.split('=', 1)
key = key.strip()
value = value.strip().strip('"').strip("'")
if key and value:
config[key] = value
except Exception as e:
print(f"⚠️ 读取配置文件失败: {e}")
return config
def get_api_key(key_name: str = "O1KEY_API_KEY") -> Optional[str]:
"""
获取 API 密钥
优先级:环境变量(推荐) > .config 文件(向后兼容)
Args:
key_name: 密钥名称,默认为 O1KEY_API_KEY
Returns:
API 密钥字符串,如果未找到则返回 None
Raises:
ValueError: 如果未找到 API 密钥
Example:
>>> api_key = get_api_key()
>>> if api_key is None:
... raise ValueError("API key not found")
"""
# 1. 优先从环境变量读取(推荐方式)
api_key = os.environ.get(key_name)
if api_key:
return api_key
# 2. 从 .config 文件读取(向后兼容,已弃用)
config = load_config()
api_key = config.get(key_name)
if api_key:
return api_key
return None
def get_api_key_or_raise(key_name: str = "O1KEY_API_KEY") -> str:
"""
获取 API 密钥,如果未找到则抛出异常
Args:
key_name: 密钥名称
Returns:
API 密钥字符串
Raises:
ValueError: 如果未找到 API 密钥
"""
api_key = get_api_key(key_name)
if not api_key:
raise ValueError("未授权!")
return api_key
+328
View File
@@ -0,0 +1,328 @@
"""
文件处理工具模块
提供文件夹图片加载、智能命名、图片配对等功能
"""
import os
import uuid
import time
from itertools import product
from pathlib import Path
from typing import List, Tuple, Optional, NamedTuple
from PIL import Image
# 支持的图片格式
SUPPORTED_IMAGE_EXTENSIONS = {'.jpg', '.jpeg', '.png', '.webp', '.bmp', '.gif'}
class ImageInfo(NamedTuple):
"""图片信息结构"""
image: Image.Image
filename: str # 不含扩展名的文件名
extension: str # 扩展名(如 .png
source_path: str # 原始文件路径
def load_images_from_folder(
folder_path: str,
recursive: bool = False
) -> List[ImageInfo]:
"""
从文件夹加载所有图片
Args:
folder_path: 文件夹路径
recursive: 是否递归加载子文件夹
Returns:
ImageInfo 列表,包含图片和元数据
Raises:
ValueError: 文件夹不存在或为空
Example:
>>> images = load_images_from_folder("D:/images")
>>> for info in images:
... print(f"{info.filename}: {info.image.size}")
"""
folder_path = folder_path.strip()
if not folder_path:
return []
path = Path(folder_path)
if not path.exists():
raise ValueError(f"文件夹不存在: {folder_path}")
if not path.is_dir():
raise ValueError(f"路径不是文件夹: {folder_path}")
images = []
# 获取文件列表
if recursive:
files = list(path.rglob("*"))
else:
files = list(path.iterdir())
# 按文件名排序,确保顺序一致
files = sorted(files, key=lambda x: x.name.lower())
for file_path in files:
if not file_path.is_file():
continue
ext = file_path.suffix.lower()
if ext not in SUPPORTED_IMAGE_EXTENSIONS:
continue
try:
img = Image.open(file_path)
img.load() # 确保图片完全加载
# 转换为 RGB 模式
if img.mode != 'RGB':
img = img.convert('RGB')
images.append(ImageInfo(
image=img,
filename=file_path.stem,
extension=ext,
source_path=str(file_path)
))
except Exception as e:
print(f"警告: 无法加载图片 {file_path}: {e}")
continue
return images
def pair_images_indexed(
*image_lists: List[ImageInfo]
) -> List[Tuple[ImageInfo, ...]]:
"""
1:1 索引配对
按索引位置配对多个图片列表,以最短列表长度为准。
Args:
*image_lists: 多个 ImageInfo 列表
Returns:
配对后的元组列表
Example:
>>> list_a = [a1, a2, a3]
>>> list_b = [b1, b2, b3]
>>> pairs = pair_images_indexed(list_a, list_b)
>>> # [(a1, b1), (a2, b2), (a3, b3)]
"""
if not image_lists:
return []
# 过滤空列表
non_empty_lists = [lst for lst in image_lists if lst]
if not non_empty_lists:
return []
# 使用 zip 进行索引配对(以最短列表为准)
return list(zip(*non_empty_lists))
def pair_images_cartesian(
*image_lists: List[ImageInfo]
) -> List[Tuple[ImageInfo, ...]]:
"""
笛卡尔积配对
生成多个图片列表的所有组合。
Args:
*image_lists: 多个 ImageInfo 列表
Returns:
配对后的元组列表
Example:
>>> list_a = [a1, a2]
>>> list_b = [b1, b2]
>>> pairs = pair_images_cartesian(list_a, list_b)
>>> # [(a1, b1), (a1, b2), (a2, b1), (a2, b2)]
"""
if not image_lists:
return []
# 过滤空列表
non_empty_lists = [lst for lst in image_lists if lst]
if not non_empty_lists:
return []
# 使用 itertools.product 生成笛卡尔积
return list(product(*non_empty_lists))
def generate_output_filename(
source_images: List[ImageInfo],
batch_index: int,
output_folder: str,
extension: str = ".png",
task_id: Optional[str] = None
) -> str:
"""
生成智能输出文件名
基于源图片文件名生成输出文件名,使用任务ID和时间戳确保并发安全。
Args:
source_images: 源图片信息列表
batch_index: 批次索引(从 0 开始)
output_folder: 输出文件夹路径
extension: 输出文件扩展名
task_id: 任务唯一标识符(用于并发场景)
Returns:
完整的输出文件路径
Example:
>>> # 单图片: hello.png -> hello_task0_12345_000.png
>>> # 多图片: hello.png + ref.png -> hello_ref_task0_12345_000.png
>>> # 并发安全:每个任务有唯一的 task_id 和时间戳
"""
# 构建基础文件名
if len(source_images) == 1:
base_name = source_images[0].filename
else:
# 多个源图片,组合文件名
names = [info.filename for info in source_images]
base_name = "_".join(names)
# 确保输出文件夹存在
output_path = Path(output_folder)
output_path.mkdir(parents=True, exist_ok=True)
# 生成唯一性标识
if task_id is None:
# 如果没有提供 task_id,使用 UUID 前8位
task_id = str(uuid.uuid4())[:8]
# 使用时间戳(毫秒级)增加唯一性
timestamp = int(time.time() * 1000) % 100000 # 精确到毫秒的后5位
# 生成文件名:基础名_任务ID_时间戳_批次索引
filename = f"{base_name}_{task_id}_{timestamp:05d}_{batch_index:03d}{extension}"
full_path = output_path / filename
# 极小概率的冲突处理
counter = 1
while full_path.exists():
filename = f"{base_name}_{task_id}_{timestamp:05d}_{batch_index:03d}_{counter}{extension}"
full_path = output_path / filename
counter += 1
return str(full_path)
def generate_batch_output_filenames(
source_images: List[ImageInfo],
count: int,
output_folder: str,
extension: str = ".png",
task_id: Optional[str] = None
) -> List[str]:
"""
批量生成输出文件名
Args:
source_images: 源图片信息列表
count: 需要生成的文件名数量
output_folder: 输出文件夹路径
extension: 输出文件扩展名
task_id: 任务唯一标识符(用于并发场景)
Returns:
输出文件路径列表
"""
filenames = []
for i in range(count):
filename = generate_output_filename(
source_images=source_images,
batch_index=i,
output_folder=output_folder,
extension=extension,
task_id=task_id
)
filenames.append(filename)
return filenames
def save_image(
image: Image.Image,
output_path: str,
quality: int = 95
) -> str:
"""
保存图片到指定路径
Args:
image: PIL Image 对象
output_path: 输出文件路径
quality: JPEG 质量(仅对 JPEG 格式有效)
Returns:
实际保存的文件路径
"""
# 确保目录存在
output_dir = Path(output_path).parent
output_dir.mkdir(parents=True, exist_ok=True)
# 根据扩展名选择保存参数
ext = Path(output_path).suffix.lower()
if ext in {'.jpg', '.jpeg'}:
# 转换为 RGBJPEG 不支持 alpha 通道)
if image.mode != 'RGB':
image = image.convert('RGB')
image.save(output_path, quality=quality)
elif ext == '.png':
image.save(output_path)
elif ext == '.webp':
image.save(output_path, quality=quality)
else:
image.save(output_path)
return output_path
def get_folder_image_count(folder_path: str) -> int:
"""
获取文件夹中的图片数量(不加载图片)
Args:
folder_path: 文件夹路径
Returns:
图片数量
"""
folder_path = folder_path.strip()
if not folder_path:
return 0
path = Path(folder_path)
if not path.exists() or not path.is_dir():
return 0
count = 0
for file_path in path.iterdir():
if file_path.is_file() and file_path.suffix.lower() in SUPPORTED_IMAGE_EXTENSIONS:
count += 1
return count
+193
View File
@@ -0,0 +1,193 @@
"""
图像处理工具模块
提供 ComfyUI Tensor 与 PIL Image 之间的转换功能
"""
import base64
from io import BytesIO
from typing import List
import numpy as np
import torch
from PIL import Image
def tensor_to_pil(tensor: torch.Tensor) -> List[Image.Image]:
"""
将 ComfyUI 的 Tensor 转换为 PIL Image 列表
Args:
tensor: 形状为 [B, H, W, C] 的张量,值范围 [0, 1]
Returns:
PIL Image 列表
Example:
>>> images = tensor_to_pil(input_tensor)
>>> for img in images:
... img.save(f"output_{i}.png")
"""
images = []
# 转换为 numpy 数组
np_images = tensor.cpu().numpy()
# 处理每张图像
for i in range(np_images.shape[0]):
img_array = np_images[i]
# 转换值范围从 [0, 1] 到 [0, 255]
img_array = (img_array * 255).astype(np.uint8)
# 创建 PIL Image
img = Image.fromarray(img_array)
images.append(img)
return images
def pil_to_tensor(images: List[Image.Image]) -> torch.Tensor:
"""
将 PIL Image 列表转换为 ComfyUI 的 Tensor
Args:
images: PIL Image 列表
Returns:
形状为 [B, H, W, C] 的张量,值范围 [0, 1]
Example:
>>> pil_images = [Image.open("test.png")]
>>> tensor = pil_to_tensor(pil_images)
>>> print(tensor.shape) # [1, H, W, 3]
"""
tensors = []
for img in images:
# 确保是 RGB 模式
if img.mode != 'RGB':
img = img.convert('RGB')
# 转换为 numpy 数组
img_array = np.array(img).astype(np.float32)
# 转换值范围从 [0, 255] 到 [0, 1]
img_array = img_array / 255.0
tensors.append(img_array)
# 堆叠为批次
batch_tensor = np.stack(tensors, axis=0)
# 转换为 torch tensor
return torch.from_numpy(batch_tensor)
def encode_image_to_base64(image: Image.Image, format: str = "PNG") -> str:
"""
将 PIL Image 编码为 base64 字符串
Args:
image: PIL Image 对象
format: 图像格式,默认 PNG
Returns:
base64 编码的字符串
Example:
>>> img = Image.open("test.png")
>>> b64_str = encode_image_to_base64(img)
"""
buffered = BytesIO()
# 转换为 RGB 模式(如果是 RGBA)
if image.mode == 'RGBA':
image = image.convert('RGB')
image.save(buffered, format=format)
img_bytes = buffered.getvalue()
return base64.b64encode(img_bytes).decode('utf-8')
def decode_base64_to_pil(base64_string: str) -> Image.Image:
"""
将 base64 字符串解码为 PIL Image
Args:
base64_string: base64 编码的图像字符串
Returns:
PIL Image 对象
Example:
>>> img = decode_base64_to_pil(b64_str)
>>> img.save("decoded.png")
"""
img_bytes = base64.b64decode(base64_string)
img = Image.open(BytesIO(img_bytes))
return img
def parse_batch_prompts(prompt: str) -> List[str]:
"""
解析批量提示词
检测单独行的 --- 分隔符,分割提示词。
如果 --- 不是单独占据一行,则返回空列表(表示单提示词模式)。
Args:
prompt: 用户输入的提示词文本
Returns:
提示词列表。如果未检测到单独行的 ---,返回空列表(表示单提示词模式)
Raises:
ValueError: 如果所有提示词都为空
Example:
>>> prompts = parse_batch_prompts("a woman\\n---\\na man")
>>> print(prompts) # ['a woman', 'a man']
>>> prompts = parse_batch_prompts("a woman --- a man")
>>> print(prompts) # [] (单提示词模式)
"""
lines = prompt.split('\n')
# 检查是否存在单独行的 ---
has_separator = False
for line in lines:
if line.strip() == '---':
has_separator = True
break
# 如果没有单独行的 ---,返回空列表(单提示词模式)
if not has_separator:
return []
# 按单独行的 --- 分割
# 先将所有单独行的 --- 替换为特殊标记
processed_lines = []
for line in lines:
if line.strip() == '---':
processed_lines.append('<<<SEPARATOR>>>')
else:
processed_lines.append(line)
# 重新组合并分割
processed_text = '\n'.join(processed_lines)
raw_prompts = processed_text.split('<<<SEPARATOR>>>')
# 过滤空提示词
filtered_prompts = []
for p in raw_prompts:
stripped = p.strip()
if stripped:
filtered_prompts.append(stripped)
# 如果所有提示词都为空,抛出错误
if not filtered_prompts:
raise ValueError("批量提示词模式下,所有提示词都为空,请至少提供一个有效的提示词")
return filtered_prompts
+83
View File
@@ -0,0 +1,83 @@
"""
更新检查工具
在插件加载时检查是否有新版本
"""
import os
import subprocess
from typing import Optional
def get_current_version() -> Optional[str]:
"""
获取当前版本号
Returns:
版本号字符串,如果读取失败返回 None
"""
version_file = os.path.join(os.path.dirname(os.path.dirname(__file__)), "version.txt")
try:
with open(version_file, 'r', encoding='utf-8') as f:
return f.read().strip()
except Exception:
return None
def check_for_updates() -> bool:
"""
检查是否有更新
Returns:
True 如果有更新,False 如果已是最新或检查失败
"""
try:
# 获取当前目录
plugin_dir = os.path.dirname(os.path.dirname(__file__))
# 检查是否是 Git 仓库
git_dir = os.path.join(plugin_dir, '.git')
if not os.path.exists(git_dir):
return False
# 执行 git fetch
subprocess.run(
['git', 'fetch', 'origin'],
cwd=plugin_dir,
capture_output=True,
timeout=10
)
# 检查本地和远程版本
local = subprocess.run(
['git', 'rev-parse', '@'],
cwd=plugin_dir,
capture_output=True,
text=True
).stdout.strip()
remote = subprocess.run(
['git', 'rev-parse', '@{u}'],
cwd=plugin_dir,
capture_output=True,
text=True
).stdout.strip()
return local != remote
except Exception:
return False
def notify_update_available():
"""通知用户有更新可用"""
current_version = get_current_version()
version_str = f" (当前版本: {current_version})" if current_version else ""
print("\n" + "="*60)
print(f"🎉 Comfyui_o1key 有新版本可用{version_str}")
print("="*60)
print("更新方法:")
print(" Windows: 双击运行 update.bat")
print(" Linux/Mac: 运行 ./update.sh")
print("或手动执行: git pull origin main")
print("="*60 + "\n")
+1
View File
@@ -0,0 +1 @@
v1.10.0
+134
View File
@@ -0,0 +1,134 @@
# Comfyui_o1key 更新说明
## 🎉 插件已支持一键自动更新!
从 v1.10.0 版本开始,插件支持自动更新功能。你只需运行更新脚本,即可轻松获取最新版本。
---
## 📦 如何更新
### Windows 用户
1. 打开文件资源管理器
2. 进入插件目录:`ComfyUI\custom_nodes\Comfyui_o1key`
3. 双击运行 `update.bat` 文件
4. 等待更新完成(通常只需几秒钟)
5. 重启 ComfyUI
### Linux/Mac 用户
打开终端,执行以下命令:
```bash
cd ComfyUI/custom_nodes/Comfyui_o1key
chmod +x update.sh # 首次运行需要添加执行权限
./update.sh
```
---
## ✨ 更新脚本功能
**自动检查更新** - 自动检测是否有新版本
**备份配置** - 自动备份和恢复 `.config` 配置文件
**拉取代码** - 自动从 GitHub 拉取最新代码
**更新依赖** - 自动更新 Python 依赖包
**显示日志** - 显示最近的更新内容
**完善提示** - 友好的中文提示和错误处理
---
## 🔔 更新检查
插件会在每次启动 ComfyUI 时自动检查是否有新版本:
- 如果发现新版本,终端会显示更新提示
- 不会影响插件加载速度
- 不会弹窗打断工作流程
- 检查失败不影响插件正常使用
**终端提示示例:**
```
============================================================
🎉 Comfyui_o1key 有新版本可用 (当前版本: v1.9.1)
============================================================
更新方法:
Windows: 双击运行 update.bat
Linux/Mac: 运行 ./update.sh
或手动执行: git pull origin main
============================================================
```
---
## 🛡️ 安全保障
- **配置安全**:更新前自动备份 `.config` 文件,更新后自动恢复
- **环境变量**:存储在系统级别的 API 密钥不受影响
- **错误处理**:更新失败不会破坏现有安装
- **回退方案**:如有问题可用 `git reset` 回退
---
## 🔧 手动更新(备用方案)
如果自动更新脚本无法使用,可以手动执行:
```bash
cd ComfyUI/custom_nodes/Comfyui_o1key
git pull origin main
pip install -r requirements.txt --upgrade
```
---
## ❓ 常见问题
### Q1: 更新会覆盖我的配置吗?
**不会。** 更新脚本会自动备份和恢复你的 `.config` 文件。环境变量中的 API 密钥也不受影响。
### Q2: 更新失败怎么办?
1. 检查网络连接是否正常
2. 确认 Git 已正确安装
3. 尝试手动更新(见上方"手动更新"部分)
4. 如有未提交的修改,先备份后执行 `git reset --hard origin/main`
### Q3: 更新后插件无法启动怎么办?
1. 检查终端错误信息
2. 重新运行依赖安装:`pip install -r requirements.txt --upgrade`
3. 确认 Python 版本 ≥ 3.7
4. 查看 [GitHub Issues](https://github.com/你的用户名/Comfyui_o1key/issues) 寻求帮助
### Q4: 可以禁用启动时的更新检查吗?
暂不支持配置禁用,但更新检查:
- 速度极快(<1 秒)
- 完全静默(无更新时不显示任何信息)
- 失败不影响插件加载
### Q5: 如何查看当前版本?
查看插件目录下的 `version.txt` 文件,或在更新时会显示当前版本号。
### Q6: 多久检查一次更新?
仅在 ComfyUI 启动时检查一次,不会在运行过程中反复检查。
---
## 📮 反馈与支持
如果在更新过程中遇到问题,请:
1. 查看终端输出的错误信息
2. 查阅 [GitHub Issues](https://github.com/你的用户名/Comfyui_o1key/issues)
3. 提交新 Issue 并附上错误信息
---
**享受自动更新带来的便利!** 🎉