1- import asyncio
21import json
32import logging
43import random
5- import requests
6- import aiohttp
74import urllib .parse
8- import urllib .request
95from typing import Optional
106
11- import websocket # NOTE: websocket-client (https://github.com/websocket-client/websocket-client)
7+ import aiohttp
128from pydantic import BaseModel
139
10+ from open_webui .env import AIOHTTP_CLIENT_SESSION_SSL
11+ from open_webui .utils .session_pool import get_session
12+
1413log = logging .getLogger (__name__ )
1514
1615default_headers = {'User-Agent' : 'Mozilla/5.0' }
1716
1817
19- def queue_prompt (prompt , client_id , base_url , api_key ):
18+ async def queue_prompt (prompt , client_id , base_url , api_key ):
2019 log .info ('queue_prompt' )
2120 p = {'prompt' : prompt , 'client_id' : client_id }
22- data = json .dumps (p ).encode ('utf-8' )
23- log .debug (f'queue_prompt data: { data } ' )
21+ log .debug (f'queue_prompt data: { p } ' )
2422 try :
25- req = urllib .request .Request (
23+ session = await get_session ()
24+ async with session .post (
2625 f'{ base_url } /prompt' ,
27- data = data ,
26+ json = p ,
2827 headers = {** default_headers , 'Authorization' : f'Bearer { api_key } ' },
29- )
30- response = urllib .request .urlopen (req ).read ()
31- return json .loads (response )
28+ ssl = AIOHTTP_CLIENT_SESSION_SSL ,
29+ ) as r :
30+ r .raise_for_status ()
31+ return await r .json ()
3232 except Exception as e :
3333 log .exception (f'Error while queuing prompt: { e } ' )
34- raise e
34+ raise
3535
3636
37- def get_image (filename , subfolder , folder_type , base_url , api_key ):
37+ async def get_image (filename , subfolder , folder_type , base_url , api_key ):
3838 log .info ('get_image' )
3939 data = {'filename' : filename , 'subfolder' : subfolder , 'type' : folder_type }
4040 url_values = urllib .parse .urlencode (data )
41- req = urllib .request .Request (
41+ session = await get_session ()
42+ async with session .get (
4243 f'{ base_url } /view?{ url_values } ' ,
4344 headers = {** default_headers , 'Authorization' : f'Bearer { api_key } ' },
44- )
45- with urllib .request .urlopen (req ) as response :
46- return response .read ()
45+ ssl = AIOHTTP_CLIENT_SESSION_SSL ,
46+ ) as r :
47+ r .raise_for_status ()
48+ return await r .read ()
4749
4850
4951def get_image_url (filename , subfolder , folder_type , base_url ):
@@ -53,32 +55,39 @@ def get_image_url(filename, subfolder, folder_type, base_url):
5355 return f'{ base_url } /view?{ url_values } '
5456
5557
56- def get_history (prompt_id , base_url , api_key ):
58+ async def get_history (prompt_id , base_url , api_key ):
5759 log .info ('get_history' )
58-
59- req = urllib . request . Request (
60+ session = await get_session ()
61+ async with session . get (
6062 f'{ base_url } /history/{ prompt_id } ' ,
6163 headers = {** default_headers , 'Authorization' : f'Bearer { api_key } ' },
62- )
63- with urllib .request .urlopen (req ) as response :
64- return json .loads (response .read ())
64+ ssl = AIOHTTP_CLIENT_SESSION_SSL ,
65+ ) as r :
66+ r .raise_for_status ()
67+ return await r .json ()
68+
6569
70+ async def _ws_get_images (ws , workflow , client_id , base_url , api_key ):
71+ """Queue a prompt and wait on *ws* for ComfyUI to finish executing it.
6672
67- def get_images (ws , workflow , client_id , base_url , api_key ):
68- prompt_id = queue_prompt (workflow , client_id , base_url , api_key )['prompt_id' ]
73+ Returns a dict of ``{'data': [{'url': ...}, ...]}``.
74+ """
75+ prompt_id = (await queue_prompt (workflow , client_id , base_url , api_key ))['prompt_id' ]
6976 output_images = []
70- while True :
71- out = ws . recv ()
72- if isinstance ( out , str ) :
73- message = json .loads (out )
77+
78+ async for msg in ws :
79+ if msg . type == aiohttp . WSMsgType . TEXT :
80+ message = json .loads (msg . data )
7481 if message ['type' ] == 'executing' :
7582 data = message ['data' ]
7683 if data ['node' ] is None and data ['prompt_id' ] == prompt_id :
7784 break # Execution is done
78- else :
79- continue # previews are binary data
85+ elif msg .type in (aiohttp .WSMsgType .CLOSED , aiohttp .WSMsgType .ERROR ):
86+ log .error (f'WebSocket closed unexpectedly: { msg .type } ' )
87+ break
88+ # binary messages (previews) are silently skipped
8089
81- history = get_history (prompt_id , base_url , api_key )[prompt_id ]
90+ history = ( await get_history (prompt_id , base_url , api_key ) )[prompt_id ]
8291 for node_id in history ['outputs' ]:
8392 node_output = history ['outputs' ][node_id ]
8493 if node_id in workflow and workflow [node_id ].get ('class_type' ) in [
@@ -105,10 +114,10 @@ async def comfyui_upload_image(image_file_item, base_url, api_key):
105114 form .add_field ('image' , file_bytes , filename = filename , content_type = mime_type )
106115 form .add_field ('type' , 'input' ) # required by ComfyUI
107116
108- async with aiohttp . ClientSession () as session :
109- async with session .post (url , data = form , headers = headers ) as resp :
110- resp .raise_for_status ()
111- return await resp .json ()
117+ session = await get_session ()
118+ async with session .post (url , data = form , headers = headers , ssl = AIOHTTP_CLIENT_SESSION_SSL ) as resp :
119+ resp .raise_for_status ()
120+ return await resp .json ()
112121
113122
114123class ComfyUINodeInput (BaseModel ):
@@ -136,11 +145,9 @@ class ComfyUICreateImageForm(BaseModel):
136145 seed : Optional [int ] = None
137146
138147
139- async def comfyui_create_image (model : str , payload : ComfyUICreateImageForm , client_id , base_url , api_key ):
140- ws_url = base_url .replace ('http://' , 'ws://' ).replace ('https://' , 'wss://' )
141- workflow = json .loads (payload .workflow .workflow )
142-
143- for node in payload .workflow .nodes :
148+ def _apply_workflow_nodes (workflow , nodes , model , payload ):
149+ """Mutate *workflow* dict in-place based on typed node definitions."""
150+ for node in nodes :
144151 if node .type :
145152 if node .type == 'model' :
146153 for node_id in node .node_ids :
@@ -151,6 +158,14 @@ async def comfyui_create_image(model: str, payload: ComfyUICreateImageForm, clie
151158 elif node .type == 'negative_prompt' :
152159 for node_id in node .node_ids :
153160 workflow [node_id ]['inputs' ][node .key if node .key else 'text' ] = payload .negative_prompt
161+ elif node .type == 'image' :
162+ if isinstance (payload .image , list ):
163+ for idx , node_id in enumerate (node .node_ids ):
164+ if idx < len (payload .image ):
165+ workflow [node_id ]['inputs' ][node .key ] = payload .image [idx ]
166+ else :
167+ for node_id in node .node_ids :
168+ workflow [node_id ]['inputs' ][node .key ] = payload .image
154169 elif node .type == 'width' :
155170 for node_id in node .node_ids :
156171 workflow [node_id ]['inputs' ][node .key if node .key else 'width' ] = payload .width
@@ -171,24 +186,31 @@ async def comfyui_create_image(model: str, payload: ComfyUICreateImageForm, clie
171186 for node_id in node .node_ids :
172187 workflow [node_id ]['inputs' ][node .key ] = node .value
173188
189+
190+ async def comfyui_create_image (model : str , payload : ComfyUICreateImageForm , client_id , base_url , api_key ):
191+ ws_url = base_url .replace ('http://' , 'ws://' ).replace ('https://' , 'wss://' )
192+ workflow = json .loads (payload .workflow .workflow )
193+ _apply_workflow_nodes (workflow , payload .workflow .nodes , model , payload )
194+
195+ headers = {'Authorization' : f'Bearer { api_key } ' }
196+ session = await get_session ()
197+
174198 try :
175- ws = websocket .WebSocket ()
176- headers = {'Authorization' : f'Bearer { api_key } ' }
177- ws .connect (f'{ ws_url } /ws?clientId={ client_id } ' , header = headers )
178- log .info ('WebSocket connection established.' )
179- except Exception as e :
199+ async with session .ws_connect (
200+ f'{ ws_url } /ws?clientId={ client_id } ' ,
201+ headers = headers ,
202+ ssl = AIOHTTP_CLIENT_SESSION_SSL ,
203+ ) as ws :
204+ log .info ('WebSocket connection established.' )
205+ log .info ('Sending workflow to WebSocket server.' )
206+ log .info (f'Workflow: { workflow } ' )
207+ images = await _ws_get_images (ws , workflow , client_id , base_url , api_key )
208+ except aiohttp .WSServerHandshakeError as e :
180209 log .exception (f'Failed to connect to WebSocket server: { e } ' )
181210 return None
182-
183- try :
184- log .info ('Sending workflow to WebSocket server.' )
185- log .info (f'Workflow: { workflow } ' )
186- images = await asyncio .to_thread (get_images , ws , workflow , client_id , base_url , api_key )
187211 except Exception as e :
188- log .exception (f'Error while receiving images: { e } ' )
189- images = None
190-
191- ws .close ()
212+ log .exception (f'Error during image generation: { e } ' )
213+ return None
192214
193215 return images
194216
@@ -209,64 +231,26 @@ class ComfyUIEditImageForm(BaseModel):
209231async def comfyui_edit_image (model : str , payload : ComfyUIEditImageForm , client_id , base_url , api_key ):
210232 ws_url = base_url .replace ('http://' , 'ws://' ).replace ('https://' , 'wss://' )
211233 workflow = json .loads (payload .workflow .workflow )
234+ _apply_workflow_nodes (workflow , payload .workflow .nodes , model , payload )
212235
213- for node in payload .workflow .nodes :
214- if node .type :
215- if node .type == 'model' :
216- for node_id in node .node_ids :
217- workflow [node_id ]['inputs' ][node .key ] = model
218- elif node .type == 'image' :
219- if isinstance (payload .image , list ):
220- # check if multiple images are provided
221- for idx , node_id in enumerate (node .node_ids ):
222- if idx < len (payload .image ):
223- workflow [node_id ]['inputs' ][node .key ] = payload .image [idx ]
224- else :
225- for node_id in node .node_ids :
226- workflow [node_id ]['inputs' ][node .key ] = payload .image
227- elif node .type == 'prompt' :
228- for node_id in node .node_ids :
229- workflow [node_id ]['inputs' ][node .key if node .key else 'text' ] = payload .prompt
230- elif node .type == 'negative_prompt' :
231- for node_id in node .node_ids :
232- workflow [node_id ]['inputs' ][node .key if node .key else 'text' ] = payload .negative_prompt
233- elif node .type == 'width' :
234- for node_id in node .node_ids :
235- workflow [node_id ]['inputs' ][node .key if node .key else 'width' ] = payload .width
236- elif node .type == 'height' :
237- for node_id in node .node_ids :
238- workflow [node_id ]['inputs' ][node .key if node .key else 'height' ] = payload .height
239- elif node .type == 'n' :
240- for node_id in node .node_ids :
241- workflow [node_id ]['inputs' ][node .key if node .key else 'batch_size' ] = payload .n
242- elif node .type == 'steps' :
243- for node_id in node .node_ids :
244- workflow [node_id ]['inputs' ][node .key if node .key else 'steps' ] = payload .steps
245- elif node .type == 'seed' :
246- seed = payload .seed if payload .seed else random .randint (0 , 1125899906842624 )
247- for node_id in node .node_ids :
248- workflow [node_id ]['inputs' ][node .key ] = seed
249- else :
250- for node_id in node .node_ids :
251- workflow [node_id ]['inputs' ][node .key ] = node .value
236+ headers = {'Authorization' : f'Bearer { api_key } ' }
237+ session = await get_session ()
252238
253239 try :
254- ws = websocket .WebSocket ()
255- headers = {'Authorization' : f'Bearer { api_key } ' }
256- ws .connect (f'{ ws_url } /ws?clientId={ client_id } ' , header = headers )
257- log .info ('WebSocket connection established.' )
258- except Exception as e :
240+ async with session .ws_connect (
241+ f'{ ws_url } /ws?clientId={ client_id } ' ,
242+ headers = headers ,
243+ ssl = AIOHTTP_CLIENT_SESSION_SSL ,
244+ ) as ws :
245+ log .info ('WebSocket connection established.' )
246+ log .info ('Sending workflow to WebSocket server.' )
247+ log .info (f'Workflow: { workflow } ' )
248+ images = await _ws_get_images (ws , workflow , client_id , base_url , api_key )
249+ except aiohttp .WSServerHandshakeError as e :
259250 log .exception (f'Failed to connect to WebSocket server: { e } ' )
260251 return None
261-
262- try :
263- log .info ('Sending workflow to WebSocket server.' )
264- log .info (f'Workflow: { workflow } ' )
265- images = await asyncio .to_thread (get_images , ws , workflow , client_id , base_url , api_key )
266252 except Exception as e :
267- log .exception (f'Error while receiving images: { e } ' )
268- images = None
269-
270- ws .close ()
253+ log .exception (f'Error during image editing: { e } ' )
254+ return None
271255
272256 return images
0 commit comments