add ollam llm api
This commit is contained in:
+3
-1
@@ -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
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -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"
|
||||
@@ -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,
|
||||
|
||||
+24
-1
@@ -28,4 +28,27 @@ def get_config_section(section_key):
|
||||
config = load_config()
|
||||
return config.get(section_key, None)
|
||||
except Exception:
|
||||
return None
|
||||
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)}")
|
||||
Reference in New Issue
Block a user