diff --git a/__init__.py b/__init__.py index 8472121..a4a564f 100644 --- a/__init__.py +++ b/__init__.py @@ -1,6 +1,7 @@ from .gemini.gemini_image_node import * from .gemini.gemini_image_preset_node import * from .ollama.ollama_vlm_node import * +from .ollama.ollama_llm_node import * from typing_extensions import override class APIExtension(ComfyExtension): @@ -9,7 +10,8 @@ class APIExtension(ComfyExtension): return [ GeminiImage, GeminiImagePreset, - OllamaVLM + OllamaVLM, + OllamaLLM ] diff --git a/ollama/ollama_llm_node.py b/ollama/ollama_llm_node.py new file mode 100644 index 0000000..4bab936 --- /dev/null +++ b/ollama/ollama_llm_node.py @@ -0,0 +1,201 @@ +import io +import os +import sys +import json +import base64 +import requests +import torch +import numpy as np +from PIL import Image +from io import BytesIO +from typing_extensions import override +from comfy_api.latest import ComfyExtension, io +from ..utils.image_utils import tensor_to_base64_string +from ..utils.config_utils import get_config_section, get_models_list + +def _load_config_credentials(): + """ + 从config.json中加载并验证API凭据 + 返回 (base_url, api_key, timeout) 元组 + """ + try: + ollama_llm_config = get_config_section('ollama-llm') + # 获取并验证base_url + if 'base_url' not in ollama_llm_config: + raise ValueError("Missing 'base_url' in ollama-vlm section") + base_url = ollama_llm_config['base_url'].strip() if isinstance(ollama_llm_config['base_url'], str) else str(ollama_llm_config['base_url']).strip() + if not base_url: + raise ValueError("base_url cannot be empty") + + # 对于Ollama,api_key是可选的 + api_key = ollama_llm_config.get('api_key', '') + api_key = api_key.strip() if isinstance(api_key, str) else str(api_key).strip() + + # 获取timeout参数,默认值为120秒 + timeout = ollama_llm_config.get('timeout', 120) + if isinstance(timeout, str): + try: + timeout = int(timeout) + except ValueError: + timeout = 120 + + return base_url, api_key, timeout + + except Exception as e: + raise ValueError(f"Failed to load Ollama VLM config section: {str(e)}") + +class OllamaLLM(io.ComfyNode): + """ + 这个节点使用 Ollama LLM 模型进行对话 + """ + # 类级别的对话历史存储,按节点实例ID存储 + _conversation_history = {} + + @classmethod + def define_schema(cls) -> io.Schema: + # 从配置文件加载模型列表 + model_options = get_models_list("ollama-llm") + default_model = model_options[0] + return io.Schema( + node_id="YCYY_Ollama_LLM_API", + display_name="Ollama LLM API", + category="YCYY/API/text", + inputs=[ + io.String.Input( + id="system_prompt", + multiline=True, + ), + io.String.Input( + id="user_prompt", + multiline=True, + ), + io.Combo.Input( + id="model", + options=model_options, + default=default_model + ), + io.Boolean.Input( + id="persist_context", + default=True, + tooltip="Persist chat context between calls (multi-turn conversation)" + ) + + ], + outputs=[ + io.String.Output( + id="Result", + display_name="Result" + ), + io.String.Output( + id="history_conversation", + display_name="History Conversation" + ) + ], + description="This node uses the Ollama LLM model for conversation." + ) + # 执行 OllamaLLM 节点 + @classmethod + def execute(cls, system_prompt, user_prompt, model, persist_context) -> io.NodeOutput: + if not user_prompt: + raise ValueError("User prompt cannot be empty") + + base_url, api_key, timeout = _load_config_credentials() + api_url = base_url+"/api/chat" + + # 生成会话标识符(基于模型和系统提示词) + session_key = f"{model}_{hash(system_prompt) if system_prompt else 'no_system'}" + + # 根据persist_context决定是否使用历史消息 + if persist_context: + # 如果会话不存在,初始化历史记录 + if session_key not in cls._conversation_history: + cls._conversation_history[session_key] = [] + # 如果有系统提示词,添加到历史记录开头 + if system_prompt: + cls._conversation_history[session_key].append({ + "role": "system", + "content": system_prompt + }) + + # 添加当前用户消息到历史 + cls._conversation_history[session_key].append({ + "role": "user", + "content": user_prompt + }) + + # 使用完整的历史消息 + messages = cls._conversation_history[session_key].copy() + else: + # 不持久化上下文,清空历史并只使用当前消息 + if session_key in cls._conversation_history: + del cls._conversation_history[session_key] + + messages = [] + if system_prompt: + messages.append({ + "role": "system", + "content": system_prompt + }) + messages.append({ + "role": "user", + "content": user_prompt + }) + + payload = { + "model": model, + "messages": messages, + "stream": False + } + + try: + if api_key: + headers = { + "Authorization": "Bearer "+api_key + } + resp = requests.post(api_url, headers=headers, json=payload, timeout=timeout) + else: + resp = requests.post(api_url, json=payload, timeout=timeout) + print(resp) + return cls._parse_response(resp, persist_context, session_key) + except Exception as e: + raise ValueError(f'The API request failed:'+{e}) + # 解析response 返回内容 + @classmethod + def _parse_response(cls, resp, persist_context, session_key): + # 检查HTTP状态码 + if resp.status_code != 200: + raise ValueError(f'API request returns an error.status_code:{resp.status_code}.error_reason:{resp.text}') + # 检查返回内容是否为空 + if not resp.text.strip(): + raise ValueError(f'The API returns an empty content') + try: + data = resp.json() + except Exception as json_exception: + # print(f"JSON解析失败:{json_exception}") + raise ValueError(f'The API returned a JSON parsing failure') + # 解析响应数据 + if "message" in data and data["message"]: + message = data.get("message", {}) + content = message.get("content", "") + + # 如果启用了上下文持久化,将助手的回复添加到历史记录 + if persist_context and session_key in cls._conversation_history: + cls._conversation_history[session_key].append({ + "role": "assistant", + "content": content + }) + + # 获取历史对话记录并转换为JSON字符串 + history_conversation = "" + if persist_context and session_key in cls._conversation_history: + history_conversation = json.dumps(cls._conversation_history[session_key], ensure_ascii=False) + else: + history_conversation = "[]" + + # 返回当前内容和历史对话JSON字符串 + return io.NodeOutput(content, history_conversation) + else: + raise ValueError(f'Content data not found') + +# 设置 web 目录,该目录中的任何 .js 文件都将被前端加载为前端扩展 +# WEB_DIRECTORY = "./somejs" diff --git a/ollama/ollama_vlm_node.py b/ollama/ollama_vlm_node.py index a3e9623..e2c921a 100644 --- a/ollama/ollama_vlm_node.py +++ b/ollama/ollama_vlm_node.py @@ -11,36 +11,9 @@ from io import BytesIO from typing_extensions import override from comfy_api.latest import ComfyExtension, io from ..utils.image_utils import tensor_to_base64_string -from ..utils.config_utils import get_config_section +from ..utils.config_utils import get_config_section,get_models_list -def _load_ollama_vlm_models(): - """ - 从config.json中加载ollama-vlm配置并获取模型列表 - """ - try: - ollama_vlm_config = get_config_section('ollama-vlm') - - # 验证配置是否存在 - if not ollama_vlm_config: - raise ValueError("Missing 'ollama-vlm' section in config file") - - # 直接获取models列表 - if 'models' not in ollama_vlm_config: - raise ValueError("Missing 'models' in ollama-vlm section") - - models = ollama_vlm_config['models'] - - # 验证models是否为列表且不为空 - if not isinstance(models, list): - raise ValueError("'models' must be a list") - - if not models: - raise ValueError("'models' list cannot be empty") - return models - except Exception as e: - raise ValueError(f"Failed to load Ollama VLM models: {str(e)}") - def _load_config_credentials(): """ 从config.json中加载并验证API凭据 @@ -85,7 +58,7 @@ class OllamaVLM(io.ComfyNode): 类型可以是 "Combo" —— 这将是一个供选择的列表。 """ # 从配置文件加载模型列表 - model_options = _load_ollama_vlm_models() + model_options = get_models_list("ollama-vlm") default_model = model_options[0] return io.Schema( node_id="YCYY_Ollama_VLM_API", @@ -96,11 +69,6 @@ class OllamaVLM(io.ComfyNode): "image", tooltip="Image used for analysis" ), - io.Combo.Input( - id="model", - options=model_options, - default=default_model - ), io.String.Input( id="system_prompt", multiline=True, diff --git a/utils/config_utils.py b/utils/config_utils.py index fe7b8a5..f7db5f1 100644 --- a/utils/config_utils.py +++ b/utils/config_utils.py @@ -28,4 +28,27 @@ def get_config_section(section_key): config = load_config() return config.get(section_key, None) except Exception: - return None \ No newline at end of file + return None +# 根据配置段 key 获取模型列表 +def get_models_list(section_key): + try: + section_config = get_config_section(section_key) + # 验证配置是否存在 + if not section_config: + raise ValueError(f"Missing {section_key} section in config file") + + # 直接获取models列表 + if 'models' not in section_config: + raise ValueError("Missing 'models' in section") + + models = section_config['models'] + + # 验证models是否为列表且不为空 + if not isinstance(models, list): + raise ValueError("'models' must be a list") + + if not models: + raise ValueError("'models' list cannot be empty") + return models + except Exception as e: + raise ValueError(f"Failed to load models: {str(e)}") \ No newline at end of file