Skip to content

Commit 5944eda

Browse files
committed
refac
1 parent 1860874 commit 5944eda

1 file changed

Lines changed: 94 additions & 110 deletions

File tree

Lines changed: 94 additions & 110 deletions
Original file line numberDiff line numberDiff line change
@@ -1,49 +1,51 @@
1-
import asyncio
21
import json
32
import logging
43
import random
5-
import requests
6-
import aiohttp
74
import urllib.parse
8-
import urllib.request
95
from typing import Optional
106

11-
import websocket # NOTE: websocket-client (https://github.com/websocket-client/websocket-client)
7+
import aiohttp
128
from 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+
1413
log = logging.getLogger(__name__)
1514

1615
default_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

4951
def 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

114123
class 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):
209231
async 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

Comments
 (0)