Update nodes.py
This commit is contained in:
@@ -1,31 +1,28 @@
|
||||
import httpx
|
||||
import json
|
||||
import mimetypes
|
||||
import asyncio
|
||||
import base64
|
||||
import os
|
||||
import uuid
|
||||
from typing import AsyncGenerator, Dict, List, Optional, Union, Any
|
||||
import aiofiles
|
||||
from typing import Generator, Dict, List, Optional, Union, Any
|
||||
import logging
|
||||
import inspect
|
||||
import io
|
||||
import threading
|
||||
import time
|
||||
|
||||
#default_system = "你是一个善于写ai画图提示词的ai助手,擅长润色提示词,描述图片,并且可以把我输入的文本和输入的图片的特征结合起来润色,不要有多余的话,直接输出描述词,结合自然语言和danbooru tags详细描述,注意千万不要忘记自然语言"
|
||||
default_system = """将danboorutag标签和图片中的特征结合(如果有)转化成自然语言,你可以将这些标签和图片(如果有)作为参考,但不要完全依赖它们,因为可能存在错误的标签。优先使用你自己的观察,并使用更合适的同义词进行描述。- 注意事项:
|
||||
- 确保场景描述的准确性。
|
||||
- 如果场景不符合上述规则或无法分析,请勿描述。保持描述的准确性。
|
||||
- 使用清晰简洁的语言进行描述。避免使用复杂或难以理解的语言。
|
||||
- 描述时尽量客观。避免使用高度主观的语言。
|
||||
- 如果场景包含NSFW或敏感内容,请勿回避。直接描述露骨内容。
|
||||
- 如果场景包含暴力、血腥、恐怖、恶心或其他敏感内容,请勿回避。直接描述。
|
||||
- 使用**英文**进行描述。
|
||||
- 想象你是一名用户试图描绘这个场景。自然地结合内容,避免结构化的描述。
|
||||
- 只返回描述。不要包含额外内容或分析。
|
||||
- 你的任务是描述图像中的每个方面、对象和互动,使得盲人在听到描述时能够在脑海中完美捕捉到图像。
|
||||
- 确保场景描述的准确性。
|
||||
- 如果场景不符合上述规则或无法分析,请勿描述。保持描述的准确性。
|
||||
- 使用清晰简洁的语言进行描述。避免使用复杂或难以理解的语言。
|
||||
- 描述时尽量客观。避免使用高度主观的语言。
|
||||
- 如果场景包含NSFW或敏感内容,请勿回避。直接描述露骨内容。
|
||||
- 如果场景包含暴力、血腥、恐怖、恶心或其他敏感内容,请勿回避。直接描述。
|
||||
- 使用**英文**进行描述。
|
||||
- 想象你是一名用户试图描绘这个场景。自然地结合内容,避免结构化的描述。
|
||||
- 只返回描述。不要包含额外内容或分析。
|
||||
- 你的任务是描述图像中的每个方面、对象和互动,使得盲人在听到描述时能够在脑海中完美捕捉到图像。
|
||||
|
||||
- !!!最重要的一点,#符号后面的tag是角色名,@符号后面的是画师名,一定要提到这两个。#和@是极其重要的特殊标签,不能删除,要将二者放在开头,作为固定开头,例如:Characters: #hoshimi miyabi. Drawn by @quan \(kurisu tina\)+xxxx(自然语言部分)。
|
||||
- !!!最重要的一点,#符号后面的tag是角色名,@符号后面的是画师名,一定要提到这两个。#和@是极其重要的特殊标签,不能删除,要将二者放在开头,作为固定开头,例如:Characters: #hoshimi miyabi. Drawn by @quan \(kurisu tina\)+xxxx(自然语言部分)。
|
||||
|
||||
特殊标签很重要
|
||||
记住了吗?"""
|
||||
@@ -44,7 +41,7 @@ except ImportError:
|
||||
|
||||
logger = logging.getLogger('OpenAI_ComfyUI')
|
||||
logger.setLevel(logging.INFO)
|
||||
from openai import AsyncOpenAI
|
||||
from openai import OpenAI
|
||||
|
||||
class OpenAIAPI:
|
||||
def __init__(
|
||||
@@ -60,17 +57,15 @@ class OpenAIAPI:
|
||||
self.apikey = apikey
|
||||
self.baseurl = baseurl if baseurl.endswith('/') else baseurl + '/'
|
||||
self.model = model
|
||||
http_client = httpx.AsyncClient(proxies=proxies, timeout=timeout) if proxies else httpx.AsyncClient(timeout=timeout)
|
||||
self.client = AsyncOpenAI(
|
||||
http_client = httpx.Client(proxies=proxies, timeout=timeout) if proxies else httpx.Client(timeout=timeout)
|
||||
self.client = OpenAI(
|
||||
api_key=apikey,
|
||||
base_url=self.baseurl,
|
||||
http_client=http_client
|
||||
)
|
||||
logger.info(f"OpenAIAPI Client Initialized: model={self.model}, base_url={self.baseurl}")
|
||||
|
||||
|
||||
async def upload_file(self, file_path: str, display_name: Optional[str] = None) -> Dict[str, Union[str, None]]:
|
||||
"""上传单个文件,使用 client.files.create,目的为 user_data"""
|
||||
def upload_file(self, file_path: str, display_name: Optional[str] = None) -> Dict[str, Union[str, None]]:
|
||||
if not os.path.exists(file_path):
|
||||
logger.error(f"文件 {file_path} 不存在")
|
||||
raise FileNotFoundError(f"文件上传失败: 路径 {file_path} 不存在")
|
||||
@@ -88,10 +83,10 @@ class OpenAIAPI:
|
||||
|
||||
try:
|
||||
logger.info(f"开始上传文件: {file_path} ({mime_type})")
|
||||
async with aiofiles.open(file_path, 'rb') as f:
|
||||
file_content = await f.read()
|
||||
with open(file_path, 'rb') as f:
|
||||
file_content = f.read()
|
||||
file_tuple = (display_name or os.path.basename(file_path), file_content, mime_type)
|
||||
file_obj = await self.client.files.create(
|
||||
file_obj = self.client.files.create(
|
||||
file=file_tuple,
|
||||
purpose="user_data"
|
||||
)
|
||||
@@ -103,8 +98,7 @@ class OpenAIAPI:
|
||||
logger.error(f"文件 {file_path} 上传失败: {type(e).__name__} - {str(e)}")
|
||||
raise RuntimeError(f"文件 {file_path} 上传失败: {type(e).__name__} - {str(e)}") from e
|
||||
|
||||
|
||||
async def _chat_api(
|
||||
def _chat_api(
|
||||
self,
|
||||
messages: List[Dict],
|
||||
stream: bool,
|
||||
@@ -112,7 +106,7 @@ class OpenAIAPI:
|
||||
topp: Optional[float] = None,
|
||||
temperature: Optional[float] = None,
|
||||
retries: int = 2
|
||||
) -> AsyncGenerator[str, None]:
|
||||
) -> Generator[str, None, None]:
|
||||
api_messages = []
|
||||
for msg in messages:
|
||||
role = msg.get("role")
|
||||
@@ -190,7 +184,6 @@ class OpenAIAPI:
|
||||
api_msg["tool_calls"] = msg["tool_calls"]
|
||||
api_messages.append(api_msg)
|
||||
|
||||
|
||||
logger.debug(f"发送给 API 的 Messages: {json.dumps(api_messages, ensure_ascii=False, indent=2)}")
|
||||
|
||||
request_params = {
|
||||
@@ -214,8 +207,8 @@ class OpenAIAPI:
|
||||
try:
|
||||
if stream:
|
||||
logger.info("发起 Stream API 请求...")
|
||||
stream_resp = await self.client.chat.completions.create(**request_params)
|
||||
async for chunk in stream_resp:
|
||||
stream_resp = self.client.chat.completions.create(**request_params)
|
||||
for chunk in stream_resp:
|
||||
if not chunk.choices:
|
||||
continue
|
||||
delta = chunk.choices[0].delta
|
||||
@@ -235,7 +228,7 @@ class OpenAIAPI:
|
||||
|
||||
else:
|
||||
logger.info(f"发起 Non-Stream API 请求 (尝试 {attempt+1}/{retries})...")
|
||||
response = await self.client.chat.completions.create(**request_params)
|
||||
response = self.client.chat.completions.create(**request_params)
|
||||
if not response.choices:
|
||||
raise RuntimeError("API 返回空 choices")
|
||||
|
||||
@@ -266,10 +259,9 @@ class OpenAIAPI:
|
||||
logger.error(f"API 调用失败 (尝试 {attempt+1}/{retries}): {type(e).__name__} - {str(e)}")
|
||||
if attempt == retries - 1:
|
||||
raise RuntimeError(f"API 调用在 {retries} 次重试后失败: {type(e).__name__} - {str(e)}") from e
|
||||
await asyncio.sleep(1.5 ** attempt)
|
||||
time.sleep(1.5 ** attempt)
|
||||
|
||||
|
||||
async def chat(
|
||||
def chat(
|
||||
self,
|
||||
messages: List[Dict[str, any]],
|
||||
stream: bool = False,
|
||||
@@ -278,7 +270,7 @@ class OpenAIAPI:
|
||||
topp: Optional[float] = None,
|
||||
temperature: Optional[float] = None,
|
||||
retries: int = 2
|
||||
) -> AsyncGenerator[str, None]:
|
||||
) -> Generator[str, None, None]:
|
||||
current_messages = list(messages)
|
||||
|
||||
if system_instruction:
|
||||
@@ -293,7 +285,7 @@ class OpenAIAPI:
|
||||
|
||||
full_response_parts = []
|
||||
try:
|
||||
async for part in self._chat_api(
|
||||
for part in self._chat_api(
|
||||
current_messages,
|
||||
stream,
|
||||
max_output_tokens,
|
||||
@@ -307,21 +299,19 @@ class OpenAIAPI:
|
||||
messages.clear()
|
||||
messages.extend(current_messages)
|
||||
|
||||
async def close_client(self):
|
||||
if self.client and hasattr(self.client, 'aclose'):
|
||||
def close_client(self):
|
||||
if self.client and hasattr(self.client, 'close'):
|
||||
logger.info("Closing httpx client...")
|
||||
await self.client.aclose()
|
||||
self.client.close()
|
||||
|
||||
API_INSTANCE_TYPE = "OPENAI_API_INSTANCE"
|
||||
CONTENT_ITEM_TYPE = "OAI_CONTENT_ITEM"
|
||||
HISTORY_TYPE = "STRING"
|
||||
|
||||
class OpenAIApiLoaderNode:
|
||||
"""加载 API 配置并创建实例"""
|
||||
def __init__(self):
|
||||
self.cached_instance : Optional[OpenAIAPI] = None
|
||||
self.cached_config_hash = None
|
||||
self.cached_loop_id = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -352,25 +342,14 @@ class OpenAIApiLoaderNode:
|
||||
config_str = f"{api_key}{model}{base_url}{proxy_http}{proxy_https}{timeout}"
|
||||
current_hash = hash(config_str)
|
||||
|
||||
try:
|
||||
current_loop = asyncio.get_running_loop()
|
||||
is_loop_ok = self.cached_loop_id == id(current_loop) and not current_loop.is_closed()
|
||||
except RuntimeError:
|
||||
current_loop = None
|
||||
is_loop_ok = False
|
||||
|
||||
if self.cached_instance and self.cached_config_hash == current_hash and is_loop_ok:
|
||||
if self.cached_instance and self.cached_config_hash == current_hash:
|
||||
logger.info("使用缓存的 OpenAI API 实例")
|
||||
return (self.cached_instance,)
|
||||
|
||||
if self.cached_instance:
|
||||
log_reason = "配置改变" if self.cached_config_hash != current_hash else "事件循环已改变或关闭"
|
||||
logger.info(f"{log_reason},关闭旧的 OpenAI API 客户端...")
|
||||
logger.info("配置改变,关闭旧的 OpenAI API 客户端...")
|
||||
try:
|
||||
if self.cached_loop_id and current_loop and not current_loop.is_closed():
|
||||
asyncio.run_coroutine_threadsafe(self.cached_instance.close_client(), current_loop)
|
||||
else:
|
||||
asyncio.run(self.cached_instance.close_client())
|
||||
self.cached_instance.close_client()
|
||||
except Exception as e:
|
||||
logger.warning(f"关闭旧客户端失败: {e}")
|
||||
self.cached_instance = None
|
||||
@@ -386,15 +365,12 @@ class OpenAIApiLoaderNode:
|
||||
)
|
||||
self.cached_instance = instance
|
||||
self.cached_config_hash = current_hash
|
||||
if current_loop:
|
||||
self.cached_loop_id = id(current_loop)
|
||||
return (instance,)
|
||||
except Exception as e:
|
||||
logger.error(f"创建 OpenAIAPI 实例失败: {e}")
|
||||
raise
|
||||
|
||||
class OpenAIImageEncoderNode:
|
||||
"""将 ComfyUI Image 编码为 Base64 内容块"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -416,8 +392,6 @@ class OpenAIImageEncoderNode:
|
||||
if image is None:
|
||||
raise ValueError("输入图像不能为空")
|
||||
|
||||
results = []
|
||||
base64_list = []
|
||||
logger.info(f"接收到 {image.shape[0]} 张图像,编码为 {format} (quality={quality}, detail={detail})")
|
||||
|
||||
try:
|
||||
@@ -451,7 +425,6 @@ class OpenAIImageEncoderNode:
|
||||
raise
|
||||
|
||||
class OpenAIFileUploaderNode:
|
||||
"""上传文件到 API 并返回文件 ID 内容块"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
input_dir = folder_paths.get_input_directory()
|
||||
@@ -471,7 +444,7 @@ class OpenAIFileUploaderNode:
|
||||
FUNCTION = "upload_file"
|
||||
CATEGORY = "OpenAI API/Content"
|
||||
|
||||
async def upload_file(self, api_instance: OpenAIAPI, file_selector: str, use_absolute_path: bool, absolute_path_override: str, display_name: str):
|
||||
def upload_file(self, api_instance: OpenAIAPI, file_selector: str, use_absolute_path: bool, absolute_path_override: str, display_name: str):
|
||||
if not api_instance:
|
||||
raise ValueError("API 实例未连接")
|
||||
|
||||
@@ -487,7 +460,7 @@ class OpenAIFileUploaderNode:
|
||||
raise FileNotFoundError(f"文件路径无效或文件不存在: {file_path}")
|
||||
|
||||
try:
|
||||
result = await api_instance.upload_file(file_path, display_name if display_name else None)
|
||||
result = api_instance.upload_file(file_path, display_name if display_name else None)
|
||||
if result and "input_file" in result and result["input_file"] and not result.get("error"):
|
||||
file_id = result["input_file"].get("file_id", "ERROR_NO_ID")
|
||||
logger.info(f"文件上传完成,file_id: {file_id}")
|
||||
@@ -497,11 +470,10 @@ class OpenAIFileUploaderNode:
|
||||
logger.error(f"文件上传节点错误: {error_msg}")
|
||||
raise RuntimeError(f"文件上传失败: {error_msg}")
|
||||
except Exception as e:
|
||||
logger.error(f"文件上传异步执行错误: {e}")
|
||||
logger.error(f"文件上传执行错误: {e}")
|
||||
raise
|
||||
|
||||
class OpenAITextBlockNode:
|
||||
"""创建纯文本内容块,可用于组合"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -518,7 +490,6 @@ class OpenAITextBlockNode:
|
||||
return ({"type": "text", "text": text},)
|
||||
|
||||
class OpenAIChatNode:
|
||||
"""执行聊天请求"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -547,13 +518,13 @@ class OpenAIChatNode:
|
||||
FUNCTION = "chat"
|
||||
CATEGORY = "OpenAI API"
|
||||
|
||||
async def chat(self, api_instance: OpenAIAPI, user_prompt: str, stream: bool, filter_reasoning: bool,
|
||||
system_prompt: str = "", history_json_in: str = "[]",
|
||||
content_part_1: Optional[Dict] = None,
|
||||
content_part_2: Optional[Dict] = None,
|
||||
content_part_3: Optional[Dict] = None,
|
||||
max_tokens: int = 1024, temperature: float = 0.7, top_p: float = 0.95, retries: int = 1,should_change: bool = False,
|
||||
):
|
||||
def chat(self, api_instance: OpenAIAPI, user_prompt: str, stream: bool, filter_reasoning: bool,
|
||||
system_prompt: str = "", history_json_in: str = "[]",
|
||||
content_part_1: Optional[Dict] = None,
|
||||
content_part_2: Optional[Dict] = None,
|
||||
content_part_3: Optional[Dict] = None,
|
||||
max_tokens: int = 1024, temperature: float = 0.7, top_p: float = 0.95, retries: int = 1, should_change: bool = False,
|
||||
):
|
||||
if not api_instance:
|
||||
raise ValueError("API 实例未连接")
|
||||
|
||||
@@ -583,12 +554,15 @@ class OpenAIChatNode:
|
||||
if not final_user_content:
|
||||
logger.warning("用户提示词和内容块均为空,跳过API调用。")
|
||||
return ("", json.dumps(messages, ensure_ascii=False, indent=2))
|
||||
|
||||
messages.append({"role": "user", "content": final_user_content})
|
||||
|
||||
async def run_chat_task(msg_list: List[Dict]):
|
||||
full_parts = []
|
||||
gen = api_instance.chat(
|
||||
messages=msg_list,
|
||||
full_parts = []
|
||||
final_text = ""
|
||||
try:
|
||||
logger.info(f"开始聊天请求 (Stream={stream})...")
|
||||
chat_generator = api_instance.chat(
|
||||
messages=messages,
|
||||
stream=stream,
|
||||
system_instruction=system_prompt if system_prompt else None,
|
||||
max_output_tokens=max_tokens,
|
||||
@@ -596,21 +570,16 @@ class OpenAIChatNode:
|
||||
topp=top_p,
|
||||
retries=retries
|
||||
)
|
||||
try:
|
||||
async for part in gen:
|
||||
if filter_reasoning and part.startswith("REASONING:"):
|
||||
logger.info(part.strip())
|
||||
continue
|
||||
full_parts.append(part)
|
||||
except Exception as e:
|
||||
logger.error(f"Chat API 异步生成器错误: {e}")
|
||||
raise
|
||||
return "".join(full_parts)
|
||||
|
||||
try:
|
||||
logger.info(f"开始聊天请求 (Stream={stream})...")
|
||||
final_text = await run_chat_task(messages)
|
||||
|
||||
for part in chat_generator:
|
||||
if filter_reasoning and part.startswith("REASONING:"):
|
||||
logger.info(part.strip())
|
||||
continue
|
||||
full_parts.append(part)
|
||||
|
||||
final_text = "".join(full_parts)
|
||||
logger.info("聊天请求结束。")
|
||||
|
||||
except Exception as e:
|
||||
final_text = f"[NODE ERROR]: {type(e).__name__} - {str(e)}"
|
||||
logger.error(f"节点执行聊天任务失败: {e}")
|
||||
|
||||
Reference in New Issue
Block a user