功能优化,app页面优化

This commit is contained in:
严浪
2025-02-09 10:32:27 +08:00
parent b1785ed5a9
commit 4402100985
26 changed files with 25137 additions and 2646 deletions
+5
View File
@@ -4,6 +4,7 @@ import os
import sys
from .lam import init, get_ext_dir
import time
from server import PromptServer
repo_dir = os.path.dirname(os.path.realpath(__file__))
sys.path.insert(0, repo_dir)
@@ -28,4 +29,8 @@ if init():
print("节点:'"+name+"'导入异常",e)
WEB_DIRECTORY = "./js"
file_directory = os.path.dirname(os.path.abspath(__file__))
PromptServer.instance.app.router.add_static("/wechatauth/static", file_directory+"/pages/static")
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS","WEB_DIRECTORY"]
+18 -5
View File
@@ -1,17 +1,30 @@
cluster:
modelPriority: false
clusterType: redis
subordinates:
- 127.0.0.1:8190
redis:
host: 127.0.0.1
port: 6379
password:
modelPriority: false #是否模型优先
base:
language: zh-CN # ja-JP ok-KR ru-RU zh-TW zh-CN en-US
api_is_used: true
appTitle: AI创作平台
appLogo: https://element-plus.sxtxhy.com/images/element-plus-logo.svg
authorIds:
- oJNTS6vtlfGKyivfY6loLLScQ3FQ
- oJNTS6vtlfGKyivfY6loLLScQ3F1
- oJNTS6vtlfGKyivfY6loLLScQ3F2
- oJNTS6vtlfGKyivfY6loLLScQ3F3
ai:
api_key: 6795fbe303878f35292f1aa14414e9a4 #智普AI开放平台注册实名认证可以免费用
model: glm-4-flash
is_tools: true #开启工具调用,默认需要配置文生图
api_key: sk-a446d27d074a45c0ac8ca1ac85742b28 #sk-a446d27d074a45c0ac8ca1ac85742b28 glm4 6795fbe303878f35292f1aa14414e9a4.kOzmNe7PEk1l6zdK sk-rttnlipmkojrsjyjrfoudbuphledqizttlqergzqxyyochfl
model: deepseek-chat # glm-4-flash deepseek-chat
ai_type: openAi #glm4 openAi
base_url: https://api.deepseek.com
is_tools: true
sys_pompt: 你是一个AI助手,能帮助我解答问题,回答内容最多不用超过200字
wechat:
+1724 -1724
View File
File diff suppressed because it is too large Load Diff
+752 -797
View File
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+78
View File
@@ -0,0 +1,78 @@
{
"title": "Material Search Engine",
"loadingModel": "Loading model...",
"scanStatus": {
"scanning": "Scanning in Progress",
"scanComplete": "Scanning Complete"
},
"statusLabels": {
"totalImages": "Total Images",
"totalVideos": "Total Videos",
"totalVideoFrames": "Total Video Frames",
"totalPexelsVideos": "Total Pexels Videos",
"scanningFiles": "Scanning Files",
"remainFiles": "Remaining Files",
"remainTime": "Estimated Remaining Time",
"scanProgress": "Scanning Progress"
},
"buttons": {
"scan": "Scan",
"cleanCache": "Clean Cache",
"logout": "Logout"
},
"searchTabs": {
"textSearch": "Text Search",
"imageSearch": "Image Search",
"textVideoSearch": "Text Video Search",
"imageVideoSearch": "Image Video Search",
"textImageSimilarity": "Text-Image Similarity Matching",
"pexelsVideos": "Pexels Videos"
},
"uploader": {
"drag": "Drag file to here, or ",
"click": "click here to upload",
"uploading": "Uploading..."
},
"formPlaceholders": {
"positiveSearch": "Search Content",
"negativeSearch": "Filter Content",
"positiveThreshold": "Search Threshold (display when similarity is above)",
"negativeThreshold": "Filter Threshold (display when similarity is below)",
"topNResults": "Show Top N Results",
"path": "Search Path (left empty means don't filter by path)",
"date": "Modify Time",
"textMatch": "Text Content (cannot be empty)",
"topnAll": "ALL"
},
"searchButtons": {
"search": "Search",
"calculateSimilarity": "Calculate Similarity"
},
"fileResults": {
"matchingProbability": "Similarity",
"matchingTimeRange": "Matching Time Range (seconds)",
"downloadVideoClip": "Download Video Clip",
"imageSearch": "Search Images",
"imageVideoSearch": "Search Videos"
},
"pexelsResults": {
"viewCount": "View Count",
"sourcePage": "Source Page"
},
"messages": {
"searchContentEmpty": "At least one search content or search path must be entered",
"textContentEmpty": "Text content cannot be empty",
"clipboardCopySuccess": "Path copied to clipboard",
"totalSearchResult": "Total search results: ",
"photos": " photos",
"videos": " videos",
"matchingSimilarityInfo": "Similarity: ",
"uploadSuccess": "Upload successful",
"imgIdNotFound": "Unable to extract image ID from img_url. This should not happen. Please report to the developer.",
"searching": "Searching..."
},
"footer": {
"description1": "This project is open-sourced on GitHub: ",
"description2": "(if you like, please star~)"
}
}
+78
View File
@@ -0,0 +1,78 @@
{
"title": "素材搜索引擎",
"loadingModel": "加载模型中……",
"scanStatus": {
"scanning": "扫描中",
"scanComplete": "扫描完成"
},
"statusLabels": {
"totalImages": "图片总数",
"totalVideos": "视频总数",
"totalVideoFrames": "视频帧总数",
"totalPexelsVideos": "pexels视频总数",
"scanningFiles": "本次扫描文件数",
"remainFiles": "待扫描文件数",
"remainTime": "预估剩余时间",
"scanProgress": "扫描进度"
},
"buttons": {
"scan": "扫描",
"cleanCache": "清除缓存",
"logout": "注销"
},
"searchTabs": {
"textSearch": "文字搜图",
"imageSearch": "以图搜图",
"textVideoSearch": "文字搜视频",
"imageVideoSearch": "以图搜视频",
"textImageSimilarity": "图文相似度匹配",
"pexelsVideos": "pexels视频"
},
"uploader": {
"drag": "将文件拖到此处,或",
"click": "点击上传",
"uploading": "上传中……"
},
"formPlaceholders": {
"positiveSearch": "搜索内容(建议使用英文逗号分隔关键词)",
"negativeSearch": "过滤内容(建议使用英文逗号分隔关键词)",
"positiveThreshold": "搜索阈值(高于该相似度才显示)",
"negativeThreshold": "过滤阈值(低于该相似度才显示)",
"topNResults": "查看前n个结果",
"path": "搜索路径(留空表示不对路径进行过滤)",
"date": "修改时间",
"textMatch": "文字内容(不能为空,建议使用英文逗号分隔关键词)",
"topnAll": "全部"
},
"searchButtons": {
"search": "搜索",
"calculateSimilarity": "计算相似度"
},
"fileResults": {
"matchingProbability": "相似度",
"matchingTimeRange": "匹配的时间段范围(秒)",
"downloadVideoClip": "下载视频片段",
"imageSearch": "以图搜图",
"imageVideoSearch": "以图搜视频"
},
"pexelsResults": {
"viewCount": "播放量",
"sourcePage": "来源页面"
},
"messages": {
"searchContentEmpty": "搜索内容或搜索路径至少需要输入一个",
"textContentEmpty": "文字内容不能为空",
"clipboardCopySuccess": "路径已复制到剪贴板",
"totalSearchResult": "共搜索出来",
"photos": "张图片",
"videos": "条视频",
"matchingSimilarityInfo": "相似度为",
"uploadSuccess": "上传成功",
"imgIdNotFound": "无法从img_url取得图片id,这个情况不应该出现,请报告给开发者",
"searching": "搜索中……"
},
"footer": {
"description1": "本项目在GitHub上开源,可以免费下载使用:",
"description2": "(求star~)"
}
}
+2 -2
View File
@@ -34,7 +34,7 @@ class SectionEnd:
if section is None:
return (images, )
if section['sectype']==1:
if len(Config().redis.keys())<=0:
if Config().cluster==None or len(Config().redis.keys())<=0:
raise Exception("redis未配置")
pool = redis.ConnectionPool(host=Config().redis['host'], port=Config().redis['port'],password=Config().redis['password'], db=0, decode_responses=True )#password="xxxxx"
rc = redis.Redis(connection_pool=pool)
@@ -54,7 +54,7 @@ class SectionEnd:
else:
fileKey=section['fileKey']
server=section['server']
if len(Config().redis.keys())<=0:
if Config().cluster==None or len(Config().redis.keys())<=0:
raise Exception("redis未配置")
pool = redis.ConnectionPool(host=Config().redis['host'], port=Config().redis['port'],password=Config().redis['password'], db=0, decode_responses=True )#password="xxxxx"
rc = redis.Redis(connection_pool=pool)
+2 -2
View File
@@ -45,7 +45,7 @@ class SectionStart:
if server=='default':
return (None,images, )
else:
if len(Config().redis.keys())<=0:
if Config().cluster==None or len(Config().redis.keys())<=0:
raise Exception("redis未配置")
pool = redis.ConnectionPool(host=Config().redis['host'], port=Config().redis['port'],password=Config().redis['password'], db=0, decode_responses=True )#password="xxxxx"
rc = redis.Redis(connection_pool=pool)
@@ -79,7 +79,7 @@ class SectionStart:
fileKey=dataObj['fileKey']
if fileKey=='':
raise Exception("文件不能为空")
if len(Config().redis.keys())<=0:
if Config().cluster==None or len(Config().redis.keys())<=0:
raise Exception("redis未配置")
pool = redis.ConnectionPool(host=Config().redis['host'], port=Config().redis['port'],password=Config().redis['password'], db=0, decode_responses=True )#password="xxxxx"
rc = redis.Redis(connection_pool=pool)
+198 -106
View File
@@ -1,5 +1,6 @@
from server import PromptServer
from server import PromptServer,send_socket_catch_exception
import nodes
import aiohttp
from aiohttp import web
import os
import json
@@ -17,11 +18,14 @@ import uuid
import folder_paths
from comfy.cli_args import args
from .src.wechat.redisSub import RedisSubscriber,run_with_reconnect,r
from .src.wechat.webSocketUtil import WebSocketClient
from threading import Thread, current_thread
from typing import List, Literal, NamedTuple, Optional
import copy
import asyncio
from zhipuai import ZhipuAI
from openai import OpenAI
import websocket
from .src.utils.chooser import ChooserMessage
# 创建一个指定长度的队列
@@ -29,17 +33,24 @@ maxsize = 10 # 队列的最大长度
client=None
userHistory={}
if len(Config().ai.keys())>0:
client = ZhipuAI(api_key=Config().ai['api_key'])
if Config().ai['ai_type']=='glm4':
client = ZhipuAI(api_key=Config().ai['api_key'])
elif Config().ai['ai_type']=='openAi':
client = OpenAI(
api_key=Config().ai['api_key'],
base_url=Config().ai['base_url'],
)
def chat_completion(userId):
tools=get_lm4_tools()
response = client.chat.completions.create(
model=Config().ai['model'], # 填写需要调用的模型名称
messages=userHistory[userId]['messages'],
tools=tools if Config().ai['is_tools'] else [],
tools=tools if Config().ai['is_tools'] else None,
tool_choice="auto" if Config().ai['is_tools'] else None, #参数设置为 “none” 来强制 API 不返回任何函数的调用。目前函数调用仅支持 auto 模式
#tool_choice={"type": "function", "function": {"name": "get_ticket_price"}}, #以强制模型生成调用get_ticket_price的参数
)
print(response)
return response
async def ai_auto_reply(msg,userId):
if client==None or len(msg.strip())==0:
@@ -53,7 +64,7 @@ async def ai_auto_reply(msg,userId):
userHistory[userId]['time'] = time.time()
try:
response=chat_completion(userId)
while response.choices[0].finish_reason=='tool_calls' and len(response.choices[0].message.tool_calls)>0:
while Config().ai['is_tools'] and response.choices[0].finish_reason=='tool_calls' and len(response.choices[0].message.tool_calls)>0:
message_handle(response.choices[0].message,userId=userId)
response = chat_completion(userId)
userHistory[userId]['messages'].append(nested_object_to_dict(response.choices[0].message))
@@ -100,7 +111,7 @@ def message_handle(message,fuctionf=None,userId=''):
for call in message.tool_calls:
message_handle(message,call,userId=userId)
def generate_image(prompt,userId,command='文生图'):
def generate_image(prompt,userId,batch_size=1,command='文生图'):
if hasattr(PromptServer.instance,'user_command') and userId in PromptServer.instance.user_command and PromptServer.instance.user_command[userId]['status']=='waiting':
msg = '您已经在队列中,请勿重复提交!'
data={'res':msg,'success':False,"res_type": "text"}
@@ -123,12 +134,15 @@ def generate_image(prompt,userId,command='文生图'):
params=Config().wechat['commands'][command]['params']
paramName=''
userData={'openId':userId,'command':command,'status':'prepare'}
userData={'openId':userId,'command':command,'status':'prepare','isAi':True}
if 'type' in Config().wechat['commands'][command]:
userData['type']=Config().wechat['commands'][command]['type']
if "prompt" in params:
userData['prompt']=prompt
if "batch_size" in params:
userData['batch_size']=batch_size
if paramName:
msg = '参数"'+paramName+'"不能为空!'
@@ -147,13 +161,7 @@ def generate_image(prompt,userId,command='文生图'):
PromptServer.instance.user_command[userId]=userData
resp=setPost(PromptServer.instance,userId)
if resp!=None:
while True:
if userId in PromptServer.instance.user_command and PromptServer.instance.user_command[userId]['status']=='waiting':
time.sleep(0.5)
else:
break
PromptServer.instance.user_command.pop(userId,None)
data = {"success": True, "res": "生成成功", "res_type": "image"}
data = {"success": True, "res": "任务下发成功", "res_type": "image"}
return data
else:
msg='非常抱歉,服务器正忙,请稍后再试!'
@@ -161,14 +169,30 @@ def generate_image(prompt,userId,command='文生图'):
return data
def addSubscribe():
sub=RedisSubscriber(Config().redis['basePath'],subscribe)
sub.run()
if Config().cluster["isMain"] and 'socket'==Config().cluster["clusterType"] and 'subordinates' in Config().cluster and len(Config().cluster["subordinates"])>0:
clients = []
for url in Config().cluster["subordinates"]:
client = WebSocketClient(url,subscribe)
client.start()
clients.append(client)
setattr(PromptServer.instance,"clients",clients)
elif 'redis'==Config().cluster["clusterType"] and r:
sub=RedisSubscriber(Config().cluster["basePath"],subscribe)
sub.run()
def subscribe(rc,msg):
message=json.loads(msg.decode())
if isinstance(rc,str):
if isinstance(msg,bytes):
send_socket_catch_exception(PromptServer.instance.sockets[PromptServer.instance.client_id].send_bytes, msg)
return
else:
message=msg
else:
message=json.loads(msg.decode())
if 'event' in message:
if message['event']=='addTask':
if Config().redis['isSection']==False:
ckptSetCount(message)
if Config().cluster["isSection"]==False:
ckptSetCount(rc,message)
prompt(PromptServer.instance,message['data'])
elif message['event']=='taskDone':
task_done(PromptServer.instance.prompt_queue,message['item_id'],message['data'])
@@ -180,7 +204,7 @@ def subscribe(rc,msg):
rc.delete(message['filename'])
elif message['event']=='sectionDone':
ChooserMessage.addMessage(**message['data'])
if Config().redis['isSection']==True and message['data']['message'] == '__cancel__':
if Config().cluster["isSection"]==True and message['data']['message'] == '__cancel__':
nodes.interrupt_processing()
elif str(message['event'])=='2':
filedata = Image.open(BytesIO(base64_to_b64decode(message['data'][1])))
@@ -191,25 +215,28 @@ def subscribe(rc,msg):
@run_with_reconnect
def ckptSetCount(message):
if 'ckptName' in message and message['ckptName']:
val=r.get('ckpt:'+Config().redis['basePath']+':'+message['ckptName'])
if val==None:
val=3
else:
val=int(val)+1
r.set('ckpt:'+Config().redis['basePath']+':'+message['ckptName'],val)
keys=r.keys('ckpt:'+Config().redis['basePath']+':*')
for key in keys:
val=int(r.get(key))
if val<=1:
r.delete(key)
def ckptSetCount(rc,message):
if isinstance(rc,str):
pass
else:
if 'ckptName' in message and message['ckptName']:
val=r.get('ckpt:'+Config().cluster["basePath"]+':'+message['ckptName'])
if val==None:
val=3
else:
r.set(key,val-1)
val=int(val)+1
r.set('ckpt:'+Config().cluster["basePath"]+':'+message['ckptName'],val)
keys=r.keys('ckpt:'+Config().cluster["basePath"]+':*')
for key in keys:
val=int(r.get(key))
if val<=1:
r.delete(key)
else:
r.set(key,val-1)
@run_with_reconnect
def sendPublish(channel,data):
if Config().redis['basePath'] == channel:
if Config().cluster["basePath"] == channel:
jsondata=json.loads(data)
if jsondata['event']=='addTask':
ckptSetCount(jsondata)
@@ -250,48 +277,48 @@ def refresh_heartbeat(prefix=''):
print('----心跳线程-----')
i=0
while True:
val=r.get(prefix+'heartbeat:'+Config().redis['basePath'])
val=r.get(prefix+'heartbeat:'+Config().cluster["basePath"])
if val==None:
val=0
r.setex(prefix+'heartbeat:'+Config().redis['basePath'], 3, val)
r.setex(prefix+'heartbeat:'+Config().cluster["basePath"], 3, val)
time.sleep(2)
if i>=30:
i=0
r.publish(Config().redis['basePath'],'{}')
r.publish(Config().cluster["basePath"],'{}')
i=i+1
@run_with_reconnect
def send_sync(self, event, data, sid=None,port=None): #继承父类的send_sync方法
if r :
if Config().redis['isMain']==False and event not in ['crystools.monitor']:
if Config().cluster and 'redis'==Config().cluster["clusterType"] and r :
if Config().cluster["isMain"]==False and event not in ['crystools.monitor']:
mainPath=r.get('mainPath')
if mainPath:
if event=='status' and Config().redis['isSection']==False:
if event=='status' and Config().cluster["isSection"]==False:
val=data['status']['exec_info']['queue_remaining']
r.setex('heartbeat:'+Config().redis['basePath'], 3, val)
r.setex('heartbeat:'+Config().cluster["basePath"], 3, val)
if event=='executed':
if 'images' in data['output']:
for img in data['output']['images']:
# imgStr=image_to_base64(img['filename'],img['type'],img['subfolder'])
# filename=base64_encode(Config().redis['basePath'])+img['filename']
# msg={'event':'sendImage','port':Config().redis['basePath'],'filename':filename,'data':imgStr
# filename=base64_encode(Config().cluster["basePath"])+img['filename']
# msg={'event':'sendImage','port':Config().cluster["basePath"],'filename':filename,'data':imgStr
# ,'type':img['type'],'subfolder':img['subfolder']}
# sendPublish(mainPath, json.dumps(msg))
# img['filename']=filename
filename=base64_encode(Config().redis['basePath'])+img['filename']
filename=base64_encode(Config().cluster["basePath"])+img['filename']
filedata=file_to_base64(img['filename'],img['type'],img['subfolder'])
r.set(filename,filedata)
msg={'event':'dowFile','port':Config().redis['basePath'],'filename':filename,'type':img['type'],'subfolder':img['subfolder']}
msg={'event':'dowFile','port':Config().cluster["basePath"],'filename':filename,'type':img['type'],'subfolder':img['subfolder']}
sendPublish(mainPath, json.dumps(msg))
img['filename']=filename
elif 'gifs' in data['output']:
for img in data['output']['gifs']:
filename=base64_encode(Config().redis['basePath'])+img['filename']
filename=base64_encode(Config().cluster["basePath"])+img['filename']
filedata=file_to_base64(img['filename'],img['type'],img['subfolder'])
r.set(filename,filedata)
msg={'event':'dowFile','port':Config().redis['basePath'],'filename':filename,'type':img['type'],'subfolder':img['subfolder']}
msg={'event':'dowFile','port':Config().cluster["basePath"],'filename':filename,'type':img['type'],'subfolder':img['subfolder']}
sendPublish(mainPath, json.dumps(msg))
img['filename']=filename
elif str(event)=='2':
@@ -302,22 +329,29 @@ def send_sync(self, event, data, sid=None,port=None): #继承父类的send_sync
filedata=base64_to_b64encode(image_data.getvalue())
datalist=[data[0],filedata,data[2]]
data=datalist
msg={'event':event,'port':Config().redis['basePath'],'data':data,'sid':sid}
msg={'event':event,'port':Config().cluster["basePath"],'data':data,'sid':sid}
sendPublish(mainPath, json.dumps(msg))
return
elif event=='status':
if port==None:
val=data['status']['exec_info']['queue_remaining']
r.setex('heartbeat:'+Config().redis['basePath'], 3, val)
r.setex('heartbeat:'+Config().cluster["basePath"], 3, val)
keys = r.keys('heartbeat:*')
queue_remaining=0
queue_remaining = PromptServer.instance.prompt_queue.get_tasks_remaining()
for key in keys:
val=r.get(key)
if val:
queue_remaining+=int(val)
data['status']['exec_info']['queue_remaining']=queue_remaining
elif hasattr(self,"clients")==True:
if event=='status':
queue_remaining=0
for c in self.clients:
queue_remaining+=c.queue_remaining
data['status']['exec_info']['queue_remaining']=queue_remaining
if hasattr(self,"clientObjPromptId")==False:
setattr(self,"clientObjPromptId",{})
if event=='execution_start':
@@ -391,14 +425,19 @@ def send_sync(self, event, data, sid=None,port=None): #继承父类的send_sync
db=DataBaseUtil()
if db.isUsable:
db.update_data('wcomplete', end_time, json.dumps(history['outputs']),data['prompt_id'])
self.user_command[sid].update({'status':'prepare','waitKey':'','seed':''.join(random.sample('123456789012345678901234567890',14))})
if 'isAi' in self.user_command[sid] and self.user_command[sid]['isAi']==True:
PromptServer.instance.user_command.pop(sid,None)
else:
self.user_command[sid].update({'status':'prepare','waitKey':'','seed':''.join(random.sample('123456789012345678901234567890',14))})
elif event == "execution_error" and hasattr(self, "user_command") and data['prompt_id'] == self.user_command[sid]['prompt_id']:
db=DataBaseUtil()
if db.isUsable:
db.delete_data(data['prompt_id'])
self.user_command[sid].update({'status':'prepare','waitKey':'','seed':''.join(random.sample('123456789012345678901234567890',14))})
if 'isAi' in self.user_command[sid] and self.user_command[sid]['isAi']==True:
PromptServer.instance.user_command.pop(sid,None)
else:
self.user_command[sid].update({'status':'prepare','waitKey':'','seed':''.join(random.sample('123456789012345678901234567890',14))})
self.loop.call_soon_threadsafe(
self.messages.put_nowait, (event, data, sid))
@@ -426,9 +465,9 @@ def task_done(self, item_id,history_result,status: Optional['PromptQueue.Executi
'status': status_dict,
}
self.history[prompt[1]].update(history_result)
if r and Config().redis['isMain']==False :
if Config().cluster and 'redis'==Config().cluster["clusterType"] and r and Config().cluster["isMain"]==False :
mainPath=r.get('mainPath')
if mainPath and Config().redis['isSection']==False:
if mainPath and Config().cluster["isSection"]==False:
sendPublish(mainPath, json.dumps({'event':'taskDone','data':self.history[prompt[1]],'item_id':prompt[1]}))
self.server.queue_updated()
@@ -478,7 +517,7 @@ def setPost(self,FromUserName):
now = time.localtime()
start_time = time.strftime("%Y-%m-%d %H:%M:%S", now)
prompt_id=str(uuid.uuid4())
if r and Config().redis['isMain']:
if Config().cluster and Config().cluster["isMain"]:
self.user_command[FromUserName]['prompt_id']=prompt_id
data=selServer(json_data,prompt_id)
if data:
@@ -525,40 +564,49 @@ def setPost(self,FromUserName):
@run_with_reconnect
def selServer(json_data,prompt_id):
json_data['prompt_id']=prompt_id
name=None
if Config().redis['modelPriority']==True :
name=getCkptName(json_data['prompt'])
if name:
ckkeys=r.keys('ckpt:*:'+name)
nport=None
for ckkey in ckkeys:
ns=ckkey.split(":")
nport=':'.join(ns[1:3])
break
if nport:
nval = r.get('heartbeat:'+nport)
if nval!=None:
sendPublish(nport, json.dumps({'event':'addTask','data':json_data,'ckptName':name}))
return {"prompt_id": prompt_id, "number": 1, "node_errors": []}
keys=r.keys('heartbeat:*')
if len(keys)>1:
nameSize={}
for key in keys:
val=r.get(key)
if val!=None :
ns=key.split(":")
print(':'.join(ns[1:]))
ckkeys=r.keys('ckpt:'+':'.join(ns[1:])+':*')
if int(val)==0 and len(ckkeys)==0:
sendPublish(':'.join(ns[1:]), json.dumps({'event':'addTask','data':json_data,'ckptName':name}))
return {"prompt_id": prompt_id, "number": 1, "node_errors": []}
else:
nameSize[':'.join(ns[1:])]=int(val)+len(ckkeys)
print('nameSize:',nameSize)
minKey=min(key for key, value in nameSize.items() if value == min(nameSize.values()))
sendPublish(minKey, json.dumps({'event':'addTask','data':json_data,'ckptName':name}))
return {"prompt_id": prompt_id, "number": 1, "node_errors": []}
if Config().cluster and 'redis'==Config().cluster["clusterType"] and r:
if Config().cluster['modelPriority']==True :
name=getCkptName(json_data['prompt'])
if name:
ckkeys=r.keys('ckpt:*:'+name)
nport=None
for ckkey in ckkeys:
ns=ckkey.split(":")
nport=':'.join(ns[1:3])
break
if nport:
nval = r.get('heartbeat:'+nport)
if nval!=None:
sendPublish(nport, json.dumps({'event':'addTask','data':json_data,'ckptName':name}))
return {"prompt_id": prompt_id, "number": 1, "node_errors": []}
keys=r.keys('heartbeat:*')
if len(keys)>1:
nameSize={}
for key in keys:
val=r.get(key)
if val!=None :
ns=key.split(":")
print(':'.join(ns[1:]))
ckkeys=r.keys('ckpt:'+':'.join(ns[1:])+':*')
if int(val)==0 and len(ckkeys)==0:
sendPublish(':'.join(ns[1:]), json.dumps({'event':'addTask','data':json_data,'ckptName':name}))
return {"prompt_id": prompt_id, "number": 1, "node_errors": []}
else:
nameSize[':'.join(ns[1:])]=int(val)+len(ckkeys)
print('nameSize:',nameSize)
minKey=min(key for key, value in nameSize.items() if value == min(nameSize.values()))
sendPublish(minKey, json.dumps({'event':'addTask','data':json_data,'ckptName':name}))
return {"prompt_id": prompt_id, "number": 1, "node_errors": []}
elif hasattr(PromptServer.instance,'clients') and len(PromptServer.instance.clients)>0:
clients=[obj for obj in PromptServer.instance.clients if obj.is_connected]
if len(clients)>0:
client = min(clients, key=lambda c: c.queue_remaining)
queue_remaining = PromptServer.instance.prompt_queue.get_tasks_remaining()
if queue_remaining>client.queue_remaining:
name=getCkptName(json_data['prompt'])
client.setSubscribe({'event':'addTask','data':json_data,'ckptName':name})
return {"prompt_id": prompt_id, "number": 1, "node_errors": []}
return None
def get_route_keys(endKey, prompt,uniqueIds):
@@ -597,7 +645,7 @@ def section_handle(json_data):
prompt[endNum]['inputs']['images']=[startNum,1]
return json_data
def trigger_on_prompt(self,json_data,isRun=True):
if isRun and r and Config().redis['isMain']:
if isRun and Config().cluster and Config().cluster["isMain"]:
prompt_id=str(uuid.uuid4())
data=selServer(json_data,prompt_id)
if data:
@@ -691,7 +739,30 @@ async def getHistorys(request):
else:
data={'msg':'openId 不能为空!','success':False}
return web.Response(text=json.dumps(data), content_type='application/json')
@PromptServer.instance.routes.post("/wechatauth/setMainServer")
async def setMainServer(request):
json_data = await request.json()
mainPath = json_data.get("mainPath")
setattr(PromptServer.instance,"mainPath",mainPath)
return web.Response(status=200)
@PromptServer.instance.routes.post("/wechatauth/setSubscribe")
async def setSubscribe(request):
json_data = await request.json()
subscribe('',json_data)
return web.Response(status=200)
@PromptServer.instance.routes.post("/wechatauth/cancelTask")
async def cancelTask(request):
post = await request.post()
prompt_id = post.get("prompt_id")
openId=post.get("openId")
delete_func = lambda a: a[1] == prompt_id
PromptServer.instance.prompt_queue.delete_queue_item(delete_func)
PromptServer.instance.user_command[openId]['status']='prepare'
return web.Response(status=200)
@PromptServer.instance.routes.post("/wechatauth/addTask")
async def addTask(request):
try:
@@ -791,12 +862,31 @@ async def getCommands(request):
return web.Response(text=json.dumps(comms), content_type='application/json')
@PromptServer.instance.routes.get("/wechatauth/app")
async def app(request):
openId=request.rel_url.query['openId']
if openId not in Config().base['authorIds']:
return web.Response(text='您没有权限访问!', content_type='text/html')
if openId in PromptServer.instance.sockets:
return web.Response(text='openId已在使用!', content_type='text/html')
basePath = folder_paths.folder_names_and_paths['custom_nodes'][0][0]
htmlPtah = os.path.join(basePath, 'ComfyUI_Lam', 'pages','app.html')
# 打开文件
with open(htmlPtah, 'r', encoding='utf-8') as file:
# 读取文件内容
html_content = file.read()
html_content = html_content.replace('{{openId}}', openId)
html_content = html_content.replace('{{appLogo}}', str(Config().base['appLogo']) if 'appLogo' in Config().base else '')
html_content = html_content.replace('{{appTitle}}', str(Config().base['appTitle']) if 'appTitle' in Config().base else '')
return web.Response(text=html_content, content_type='text/html')
@PromptServer.instance.routes.get("/wechatauth/app2")
async def app(request):
if "openId" in request.rel_url.query:
openId=request.rel_url.query['openId']
openId=base64_decode(openId)
basePath = folder_paths.folder_names_and_paths['custom_nodes'][0][0]
htmlPtah = os.path.join(basePath, 'ComfyUI_Lam', 'pages','app.html')
htmlPtah = os.path.join(basePath, 'ComfyUI_Lam', 'pages','app2.html')
# 打开文件
with open(htmlPtah, 'r', encoding='utf-8') as file:
# 读取文件内容
@@ -933,7 +1023,7 @@ async def handleMessagePost(request):
data=otherName.split('加')
openId=base64_decode(data[0])
tount=int(data[1])
if openId == 'config':
if data[0] == 'config':
Config().reload()
msg='配置文件已更新'
elif openId and tount>0:
@@ -946,7 +1036,6 @@ async def handleMessagePost(request):
else:
msg='用户不存在'
else:
msg='数据库连接失败'
else:
@@ -1032,19 +1121,22 @@ PromptServer.instance.old_trigger_on_prompt=PromptServer.instance.trigger_on_pro
PromptServer.instance.trigger_on_prompt=types.MethodType(trigger_on_prompt,PromptServer.instance)
if hasattr(PromptServer.instance,"displayName")==False:
setattr(PromptServer.instance,"displayName",NODE_LANGEUAGE_DISPLAY_NAME_MAPPINGS)
if r: #添加订阅消息
if Config().redis['isMain']:
r.set('mainPath',Config().redis['basePath'])
keys=r.keys('ckpt:'+Config().redis['basePath']+':*')
for key in keys:
r.delete(key)
prefix=''
if Config().redis['isSection']:
prefix='section'
Thread(target=refresh_heartbeat,daemon=True, args=(prefix,)).start()
if Config().cluster:
if'redis'==Config().cluster["clusterType"] and r: #添加订阅消息
if Config().cluster["isMain"]:
r.set('mainPath',Config().cluster["basePath"])
keys=r.keys('ckpt:'+Config().cluster["basePath"]+':*')
for key in keys:
r.delete(key)
prefix=''
if Config().cluster["isSection"]:
prefix='section'
Thread(target=refresh_heartbeat,daemon=True, args=(prefix,)).start()
Thread(target=addSubscribe,daemon=True, args=()).start()
if PromptServer.instance.prompt_queue:
PromptServer.instance.prompt_queue.task_done=types.MethodType(task_done,PromptServer.instance.prompt_queue)
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
+6 -1
View File
@@ -65,8 +65,13 @@ def base64_encode(text):
def base64_decode(encoded_text):
'''解密'''
decoded_text = base64.b64decode(encoded_text).decode('utf-8')
try:
decoded_text = base64.b64decode(encoded_text).decode('utf-8')
except Exception as e:
logging.error(f"base64_decode error: {e}")
decoded_text=''
return decoded_text
"""这是一个处理客户发送信息的文件"""
+7 -6
View File
@@ -60,16 +60,17 @@ class Config(object):
self.wechat = yconfig.get("wechat", {})
self.base = yconfig.get("base", {})
self.ai = yconfig.get("ai", {})
self.redis = yconfig.get("redis", {})
if "cluster" in args and args.cluster:
self.redis = yconfig.get("redis", {})
self.redis["isSection"] = args.isSection
self.cluster = yconfig.get("cluster", {})
self.cluster["isSection"] = args.isSection
if args.isSection:
self.redis["isMain"] = False
self.cluster["isMain"] = False
else:
self.redis["isMain"] = args.isMain
self.redis["basePath"] = args.basePath+":"+str(args.port)
self.cluster["isMain"] = args.isMain
self.cluster["basePath"] = args.basePath+":"+str(args.port)
else:
self.redis = {}
self.cluster = None
#self.EMAIL= yconfig.get("email", {})
#self.OPENAI= yconfig.get("openai", {})
#self.GLM4= yconfig.get("glm4", {})
+1 -1
View File
@@ -7,7 +7,7 @@ import time
r=None
def connect_redis():
global r
if len(Config().redis.keys())>0:
if Config().cluster and 'redis'==Config().cluster["clusterType"] and len(Config().redis.keys())>0:
print('连接redis')
# 尝试连接Redis
try:
+1 -1
View File
@@ -152,7 +152,7 @@ def resetting_chat_record() -> dict:
@register_tool
def generate_image(prompt: Annotated[str, '要生成图片的英文提示词', True]) -> dict:
def generate_image(prompt: Annotated[str, '要生成图片的英文提示词', True],batch_size:Annotated[int, '图片数量', True]) -> dict:
'''
生成图片`prompt`英文提示词
'''
+97
View File
@@ -0,0 +1,97 @@
import websocket
import uuid
import threading
import time
import json
import requests
import urllib
import folder_paths
import os
from .config import Config
class WebSocketClient(threading.Thread):
def __init__(self, server_address, messageFunc=None):
super().__init__()
self.server_address = server_address
self.client_id=str(uuid.uuid4())
self.ws = websocket.WebSocket()
self.messageFunc = messageFunc
self.queue_remaining=0
self.cliDict={}
self.is_connected=False
def handle_message(self, data):
if self.messageFunc != None:
self.messageFunc(self.server_address, data)
def setMainServer(self):
p = {"mainPath": self.client_id}
data = json.dumps(p).encode('utf-8')
req = requests.post("http://{}/wechatauth/setMainServer".format(self.server_address), data=data)
print("返回状态码:",req.status_code)
if req.status_code==200:
self.is_connected=True
def setSubscribe(self,json_data):
if json_data['event']=='addTask':
self.cliDict[json_data['data']['prompt_id']]=json_data['data']['client_id']
json_data['data']['client_id']=self.client_id
data = json.dumps(json_data).encode('utf-8')
req = requests.post("http://{}/wechatauth/setSubscribe".format(self.server_address), data=data)
print("返回状态码:",req.status_code)
return req.status_code==200
def output(self,node_output):
if 'images' in node_output:
for image in node_output['images']:
image['filename']=self.dowFile(image)
elif 'gifs' in node_output:
for image in node_output['gifs']:
image['filename']=self.dowFile(image)
def dowFile(self,data):
if data['type']=="temp":
output_dir = folder_paths.get_temp_directory()
else:
output_dir = folder_paths.get_output_directory()
if 'subfolder' in data and data['subfolder']:
file_path=os.path.join(output_dir,data['subfolder'],self.client_id+data['filename'])
else:
file_path=os.path.join(output_dir,self.client_id+data['filename'])
url_values = urllib.parse.urlencode(data)
with requests.get("http://{}/view?{}".format(self.server_address, url_values)) as response:
open(file_path,'wb').write(response.content)
return self.client_id+data['filename']
def restart(self):
url="ws://{}/ws?clientId={}".format(self.server_address, self.client_id)
self.ws.connect(url)
self.setMainServer()
while True:
out = self.ws.recv()
if isinstance(out, str):
message = json.loads(out)
if 'type' in message:
message['event']=message['type']
if 'event' in message:
if message['event'] in ['crystools.monitor']:
continue
if message['event'] == 'status':
self.queue_remaining=message['data']['status']['exec_info']['queue_remaining']
elif message['event']=='executed':
self.output(message['data']['output'])
if 'data' in message and 'promptId' in message['data']:
message['sid']=self.cliDict[message['data']['promptId']]
else:
message['sid']=None
message['port']=self.server_address
self.handle_message(message)
def run(self):
while True:
try:
self.restart()
except Exception as e:
self.is_connected=False
time.sleep(2)
print(f'websocket异常:{e},稍后重连')
+2 -1
View File
@@ -20,4 +20,5 @@ face_alignment
pyzbar
redis
zhipuai
opencv-contrib-python
opencv-contrib-python
websocket-client