feat: sync latest local version as authoritative codebase

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

Co-Authored-By: Claude Sonnet 4.5 <[email protected]>
This commit is contained in:
o1key
2026-04-03 16:18:45 +08:00
co-authored by Claude Sonnet 4.5
commit 9ab209b2b7
46 changed files with 16757 additions and 0 deletions
+37
View File
@@ -0,0 +1,37 @@
"""
工具模块
包含图像处理、配置管理、文件处理等通用工具函数
"""
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_timestamp_filename,
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_timestamp_filename',
'save_image',
'get_folder_image_count'
]
+117
View File
@@ -0,0 +1,117 @@
"""
配置管理模块
处理环境变量和 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")
# ============ API 基础配置 ============
# 所有 API 客户端的统一基础 URL
# 可通过环境变量 O1KEY_API_BASE_URL 覆盖
DEFAULT_API_BASE_URL = "https://vip.o1key.com"
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
"""
config = load_config()
return config.get(key_name)
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
def get_api_base_url() -> str:
"""
获取 API 基础 URL
从 .config 文件读取,如果未配置则使用默认值
Returns:
API 基础 URL 字符串
"""
config = load_config()
base_url = config.get("O1KEY_API_BASE_URL")
if base_url:
return base_url.rstrip('/')
return DEFAULT_API_BASE_URL
+40
View File
@@ -0,0 +1,40 @@
"""
文件数据类型定义
用于在 ComfyUI 节点间传递文件数据
"""
from typing import NamedTuple
class FileData(NamedTuple):
"""
文件数据类型,用于在节点间传递
Attributes:
path: 文件完整路径
filename: 文件名(不含扩展名)
extension: 文件扩展名(如 .pdf
mime_type: MIME 类型
data: Base64 编码的文件内容
size: 文件大小(字节)
"""
path: str
filename: str
extension: str
mime_type: str
data: str
size: int
# 支持的文档 MIME 类型映射
DOCUMENT_MIME_TYPES = {
".pdf": "application/pdf",
".txt": "text/plain"
}
# 文件大小限制(字节)
FILE_SIZE_LIMITS = {
".pdf": 50 * 1024 * 1024, # 50MB (Gemini API 官方限制)
".txt": 20 * 1024 * 1024 # 20MB (保守限制)
}
+375
View File
@@ -0,0 +1,375 @@
"""
文件处理工具模块
提供文件夹图片加载、智能命名、图片配对等功能
"""
import os
import re
import uuid
import time
import random
from datetime import datetime
from itertools import product
from pathlib import Path
from typing import List, Tuple, Optional, NamedTuple
from PIL import Image
def _get_server_port() -> Optional[int]:
"""获取当前 ComfyUI 实例的端口号,失败返回 None"""
try:
import comfy.cli_args
port = getattr(comfy.cli_args.args, 'port', None) or getattr(comfy.cli_args, 'server_port', None) or getattr(comfy.cli_args, 'port', None)
if port is not None:
return int(port)
except Exception:
pass
# 备用:从 listen 环境变量或命令行参数尝试
try:
import sys
for arg in sys.argv:
if '--port' in arg or '--listen-port' in arg:
parts = arg.split('=')
if len(parts) == 2:
return int(parts[1].strip())
elif arg in ('--port', '--listen-port'):
idx = sys.argv.index(arg)
if idx + 1 < len(sys.argv):
return int(sys.argv[idx + 1])
except Exception:
pass
return None
def _get_port_suffix() -> str:
"""
返回非默认端口的后缀字符串(如 "_8189"),默认端口 8188 或获取失败时返回空字符串。
"""
try:
port = _get_server_port()
if port is not None and port != 8188:
return f"_{port}"
except Exception:
pass
return ""
# 支持的图片格式
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_by_name(
*image_lists: List[ImageInfo]
) -> List[Tuple[ImageInfo, ...]]:
"""
按文件名配对(同名匹配)
取所有文件夹中文件名(不含扩展名)的交集,按文件名字母升序排列后配对。
只有在所有文件夹中都存在同名文件,该文件名才会被纳入配对。
扩展名不同的文件(如 1.jpg 与 1.png)视为同名。
Args:
*image_lists: 多个 ImageInfo 列表
Returns:
配对后的元组列表,按文件名字母升序排列
Raises:
ValueError: 所有文件夹之间没有任何相同文件名时抛出
Example:
>>> list_a = [ImageInfo(filename="1", ...), ImageInfo(filename="2", ...)]
>>> list_b = [ImageInfo(filename="1", ...), ImageInfo(filename="3", ...)]
>>> pairs = pair_images_by_name(list_a, list_b)
>>> # [(list_a[0], list_b[0])] # 只有 "1" 匹配
"""
if not image_lists:
return []
non_empty_lists = [lst for lst in image_lists if lst]
if not non_empty_lists:
return []
# 单文件夹直接返回(无需配对)
if len(non_empty_lists) == 1:
return [(img,) for img in non_empty_lists[0]]
# 为每个文件夹建立 filenamestem-> ImageInfo 的映射
name_maps = [
{img.filename: img for img in lst}
for lst in non_empty_lists
]
# 取所有文件夹文件名的交集
common_names = set(name_maps[0].keys())
for nm in name_maps[1:]:
common_names &= set(nm.keys())
if not common_names:
# 收集各文件夹的文件名示例,帮助用户排查问题
folder_samples = []
for i, nm in enumerate(name_maps):
sample = sorted(nm.keys())[:3]
sample_str = "".join(f'"{n}"' for n in sample)
folder_samples.append(f"文件夹{i + 1}{sample_str}")
samples_info = "\n".join(folder_samples)
raise ValueError(
f"所有文件夹中没有找到任何同名图片,无法进行配对!\n"
f"请确保各文件夹内存在文件名相同的图片后重试。\n"
f"(文件名比较不含扩展名,例如「1.jpg」与「1.png」视为同名)\n\n"
f"各文件夹当前文件名示例:\n{samples_info}"
)
# 按文件名字母升序排列,保证顺序稳定
sorted_names = sorted(common_names, key=lambda x: x.lower())
return [
tuple(nm[name] for nm in name_maps)
for name in sorted_names
]
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_timestamp_filename(output_folder: str, prefix: str = "", extension: str = ".png", port_suffix: str = "") -> str:
"""
生成基于时间戳的文件名,确保按文件名排序 = 按生成时间排序。
格式:{prefix}{HHMMSS_YYYYMMDD_mmm}{port_suffix}{extension}
例如:161700_20260322_001.png 或 去除ai_161700_20260322_001.png
Args:
output_folder: 输出目录
prefix: 文件名前缀(如 "去除ai_"
extension: 文件扩展名(如 ".png"
port_suffix: 端口后缀(如 "_8189"),为空时自动获取
Returns:
完整文件路径
"""
Path(output_folder).mkdir(parents=True, exist_ok=True)
if not port_suffix:
port_suffix = _get_port_suffix()
date_part = datetime.now().strftime("%Y%m%d")
time_part = datetime.now().strftime("%H%M%S")
ms = random.randint(0, 999)
while True:
filename = f"{prefix}{time_part}_{date_part}_{ms:03d}{port_suffix}{extension}"
full_path = Path(output_folder) / filename
if not full_path.exists():
return str(full_path)
ms = (ms + 1) % 1000
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 == '.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
+86
View File
@@ -0,0 +1,86 @@
"""
更新检查工具
在插件加载时检查是否有新版本
"""
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(禁止弹出认证弹框,失败时静默处理)
env = os.environ.copy()
env['GIT_TERMINAL_PROMPT'] = '0'
subprocess.run(
['git', 'fetch', 'origin'],
cwd=plugin_dir,
capture_output=True,
timeout=10,
env=env
)
# 检查本地和远程版本
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")