Changed HTTP Server + Added WebSockets

Moved the existing API endpoints to use aoihttp and added websocket notifications
This commit is contained in:
pythongosssss 2023-02-12 15:53:48 +00:00 committed by GitHub
parent f542f248f1
commit 5d14e9b959
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23

206
main.py
View File

@ -5,6 +5,7 @@ import json
import threading import threading
import heapq import heapq
import traceback import traceback
import asyncio
if __name__ == "__main__": if __name__ == "__main__":
if '--help' in sys.argv: if '--help' in sys.argv:
@ -25,7 +26,6 @@ if '--dont-upcast-attention' in sys.argv:
os.environ['ATTN_PRECISION'] = "fp16" os.environ['ATTN_PRECISION'] = "fp16"
import torch import torch
import nodes import nodes
def get_input_data(inputs, class_def, outputs={}, prompt={}, extra_data={}): def get_input_data(inputs, class_def, outputs={}, prompt={}, extra_data={}):
@ -286,16 +286,19 @@ def prompt_worker(q):
q.task_done(item_id) q.task_done(item_id)
class PromptQueue: class PromptQueue:
def __init__(self): def __init__(self, socket_handler):
self.socket_handler = socket_handler
self.mutex = threading.RLock() self.mutex = threading.RLock()
self.not_empty = threading.Condition(self.mutex) self.not_empty = threading.Condition(self.mutex)
self.task_counter = 0 self.task_counter = 0
self.queue = [] self.queue = []
self.currently_running = {} self.currently_running = {}
socket_handler.prompt_queue = self
def put(self, item): def put(self, item):
with self.mutex: with self.mutex:
heapq.heappush(self.queue, item) heapq.heappush(self.queue, item)
self.socket_handler.queue_updated(self)
self.not_empty.notify() self.not_empty.notify()
def get(self): def get(self):
@ -306,11 +309,13 @@ class PromptQueue:
i = self.task_counter i = self.task_counter
self.currently_running[i] = copy.deepcopy(item) self.currently_running[i] = copy.deepcopy(item)
self.task_counter += 1 self.task_counter += 1
self.socket_handler.queue_updated(self)
return (item, i) return (item, i)
def task_done(self, item_id): def task_done(self, item_id):
with self.mutex: with self.mutex:
self.currently_running.pop(item_id) self.currently_running.pop(item_id)
self.socket_handler.queue_updated(self)
def get_current_queue(self): def get_current_queue(self):
with self.mutex: with self.mutex:
@ -326,6 +331,7 @@ class PromptQueue:
def wipe_queue(self): def wipe_queue(self):
with self.mutex: with self.mutex:
self.queue = [] self.queue = []
self.socket_handler.queue_updated(self)
def delete_queue_item(self, function): def delete_queue_item(self, function):
with self.mutex: with self.mutex:
@ -336,35 +342,82 @@ class PromptQueue:
else: else:
self.queue.pop(x) self.queue.pop(x)
heapq.heapify(self.queue) heapq.heapify(self.queue)
self.socket_handler.queue_updated(self)
return True return True
return False return False
from http.server import BaseHTTPRequestHandler, HTTPServer import aiohttp
from aiohttp import web
class PromptServer(BaseHTTPRequestHandler): def get_queue_info(prompt_queue):
def _set_headers(self, code=200, ct='text/html'): prompt_info = {}
self.send_response(code) exec_info = {}
self.send_header('Content-type', ct) exec_info['queue_remaining'] = prompt_queue.get_tasks_remaining()
self.end_headers() prompt_info['exec_info'] = exec_info
def log_message(self, format, *args): return prompt_info
pass
def do_GET(self): class SocketHandler():
if self.path == "/prompt": def __init__(self, loop):
self._set_headers(ct='application/json') self.connected = set()
prompt_info = {} self.messages = asyncio.Queue()
exec_info = {} self.loop = loop
exec_info['queue_remaining'] = self.server.prompt_queue.get_tasks_remaining()
prompt_info['exec_info'] = exec_info async def publish_loop(self):
self.wfile.write(json.dumps(prompt_info).encode('utf-8')) while True:
elif self.path == "/queue": msg = await self.messages.get()
self._set_headers(ct='application/json') await self.send(msg)
queue_info = {}
current_queue = self.server.prompt_queue.get_current_queue() def queue_updated(self, queue):
queue_info['queue_running'] = current_queue[0] # This is called by the queue processing thread so we need to make it thread safe
queue_info['queue_pending'] = current_queue[1] loop.call_soon_threadsafe(self.messages.put_nowait, { 'type': 'status', 'status': get_queue_info(queue) })
self.wfile.write(json.dumps(queue_info).encode('utf-8'))
elif self.path == "/object_info": async def send(self, message, socket = None):
self._set_headers(ct='application/json') if isinstance(message, str) == False:
message = json.dumps(message)
if socket is None:
for ws in self.connected:
await ws.send_str(message)
else:
await socket.send_str(message)
async def process(self, request):
ws = web.WebSocketResponse()
await ws.prepare(request)
self.connected.add(ws)
try:
# Send initial state to the new client
await self.send({ 'type': 'status', 'status': get_queue_info(self.prompt_queue) }, ws)
async for msg in ws:
if msg.type == aiohttp.WSMsgType.ERROR:
print('ws connection closed with exception %s' % ws.exception())
finally:
self.connected.remove(ws)
return ws
class PromptServer():
def __init__(self, prompt_queue, socket_handler):
self.prompt_queue = prompt_queue
self.socket_handler = socket_handler
self.number = 0
self.app = web.Application()
routes = web.RouteTableDef()
@routes.get('/ws')
async def websocket_handler(request):
return await self.socket_handler.process(request)
@routes.get("/")
async def get_root(request):
return aiohttp.web.HTTPFound('/index.html')
@routes.get("/prompt")
async def get_prompt(request):
return web.json_response(get_queue_info(self.prompt_queue))
@routes.get("/object_info")
async def get_object_info(request):
out = {} out = {}
for x in nodes.NODE_CLASS_MAPPINGS: for x in nodes.NODE_CLASS_MAPPINGS:
obj_class = nodes.NODE_CLASS_MAPPINGS[x] obj_class = nodes.NODE_CLASS_MAPPINGS[x]
@ -377,40 +430,32 @@ class PromptServer(BaseHTTPRequestHandler):
if hasattr(obj_class, 'CATEGORY'): if hasattr(obj_class, 'CATEGORY'):
info['category'] = obj_class.CATEGORY info['category'] = obj_class.CATEGORY
out[x] = info out[x] = info
self.wfile.write(json.dumps(out).encode('utf-8')) return web.json_response(out)
elif self.path[1:] in os.listdir(self.server.server_dir):
if self.path[1:].endswith('.css'):
self._set_headers(ct='text/css')
elif self.path[1:].endswith('.js'):
self._set_headers(ct='text/javascript')
else:
self._set_headers()
with open(os.path.join(self.server.server_dir, self.path[1:]), "rb") as f:
self.wfile.write(f.read())
else:
self._set_headers()
with open(os.path.join(self.server.server_dir, "index.html"), "rb") as f:
self.wfile.write(f.read())
def do_HEAD(self): @routes.get("/queue")
self._set_headers() async def get_queue(request):
queue_info = {}
current_queue = self.prompt_queue.get_current_queue()
queue_info['queue_running'] = current_queue[0]
queue_info['queue_pending'] = current_queue[1]
return web.json_response(queue_info)
def do_POST(self): @routes.post("/prompt")
resp_code = 200 async def post_prompt(request):
out_string = ""
if self.path == "/prompt":
print("got prompt") print("got prompt")
data_string = self.rfile.read(int(self.headers['Content-Length'])) resp_code = 200
json_data = json.loads(data_string) out_string = ""
json_data = await request.json()
if "number" in json_data: if "number" in json_data:
number = float(json_data['number']) number = float(json_data['number'])
else: else:
number = self.server.number number = self.number
if "front" in json_data: if "front" in json_data:
if json_data['front']: if json_data['front']:
number = -number number = -number
self.server.number += 1 self.number += 1
if "prompt" in json_data: if "prompt" in json_data:
prompt = json_data["prompt"] prompt = json_data["prompt"]
valid = validate_prompt(prompt) valid = validate_prompt(prompt)
@ -418,46 +463,54 @@ class PromptServer(BaseHTTPRequestHandler):
if "extra_data" in json_data: if "extra_data" in json_data:
extra_data = json_data["extra_data"] extra_data = json_data["extra_data"]
if valid[0]: if valid[0]:
self.server.prompt_queue.put((number, id(prompt), prompt, extra_data)) self.prompt_queue.put((number, id(prompt), prompt, extra_data))
else: else:
resp_code = 400 resp_code = 400
out_string = valid[1] out_string = valid[1]
print("invalid prompt:", valid[1]) print("invalid prompt:", valid[1])
elif self.path == "/queue":
data_string = self.rfile.read(int(self.headers['Content-Length'])) return web.Response(body=out_string, status=resp_code)
json_data = json.loads(data_string)
@routes.post("/queue")
async def post_queue(request):
json_data = await request.json()
if "clear" in json_data: if "clear" in json_data:
if json_data["clear"]: if json_data["clear"]:
self.server.prompt_queue.wipe_queue() self.prompt_queue.wipe_queue()
if "delete" in json_data: if "delete" in json_data:
to_delete = json_data['delete'] to_delete = json_data['delete']
for id_to_delete in to_delete: for id_to_delete in to_delete:
delete_func = lambda a: a[1] == int(id_to_delete) delete_func = lambda a: a[1] == int(id_to_delete)
self.server.prompt_queue.delete_queue_item(delete_func) self.prompt_queue.delete_queue_item(delete_func)
self._set_headers(code=resp_code) return web.Response(status=200)
self.end_headers()
self.wfile.write(out_string.encode('utf8'))
return
self.app.add_routes(routes)
self.app.add_routes([
web.static('/', os.path.join(os.path.dirname(os.path.realpath(__file__)), "webshit")),
])
def run(prompt_queue, address='', port=8188): async def start_server(server, address, port):
server_address = (address, port) runner = web.AppRunner(server.app)
httpd = HTTPServer(server_address, PromptServer) await runner.setup()
httpd.server_dir = os.path.join(os.path.dirname(os.path.realpath(__file__)), "webshit") site = web.TCPSite(runner, address, port)
httpd.prompt_queue = prompt_queue await site.start()
httpd.number = 0
if server_address[0] == '': if address == '':
addr = '0.0.0.0' address = '0.0.0.0'
else:
addr = server_address[0]
print("Starting server\n") print("Starting server\n")
print("To see the GUI go to: http://{}:{}".format(addr, server_address[1])) print("To see the GUI go to: http://{}:{}".format(address, port))
httpd.serve_forever()
async def run(prompt_queue, socket_handler, address='', port=8188):
server = PromptServer(prompt_queue, socket_handler)
await asyncio.gather(start_server(server, address, port), socket_handler.publish_loop())
if __name__ == "__main__": if __name__ == "__main__":
q = PromptQueue() loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
socket_handler = SocketHandler(loop)
q = PromptQueue(socket_handler)
threading.Thread(target=prompt_worker, daemon=True, args=(q,)).start() threading.Thread(target=prompt_worker, daemon=True, args=(q,)).start()
if '--listen' in sys.argv: if '--listen' in sys.argv:
address = '0.0.0.0' address = '0.0.0.0'
@ -471,6 +524,9 @@ if __name__ == "__main__":
except: except:
pass pass
run(q, address=address, port=port) try:
loop.run_until_complete(run(q, socket_handler, address=address, port=port))
except KeyboardInterrupt:
pass