Coverage for open_webui/utils/images/comfyui.py: 19%

178 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1import logging 

2import random 

3import urllib.parse 

4from typing import Optional 

5 

6import aiohttp 

7from open_webui.env import AIOHTTP_CLIENT_SESSION_SSL 

8from open_webui.utils.json_codec import JSONCodec 

9from open_webui.utils.session_pool import get_session 

10from pydantic import BaseModel 

11 

12log = logging.getLogger(__name__) 

13 

14default_headers = {'User-Agent': 'Mozilla/5.0'} 

15 

16 

17async def queue_prompt(prompt, client_id, base_url, api_key): 

18 log.info('queue_prompt') 

19 p = {'prompt': prompt, 'client_id': client_id} 

20 log.debug('queue_prompt data: %s', p) 

21 try: 

22 session = await get_session() 

23 async with session.post( 

24 f'{base_url}/prompt', 

25 json=p, 

26 headers={**default_headers, 'Authorization': f'Bearer {api_key}'}, 

27 ssl=AIOHTTP_CLIENT_SESSION_SSL, 

28 ) as r: 

29 r.raise_for_status() 

30 return await r.json() 

31 except Exception as e: 

32 log.exception(f'Error while queuing prompt: {e}') 

33 raise 

34 

35 

36async def get_image(filename, subfolder, folder_type, base_url, api_key): 

37 log.info('get_image') 

38 data = {'filename': filename, 'subfolder': subfolder, 'type': folder_type} 

39 url_values = urllib.parse.urlencode(data) 

40 session = await get_session() 

41 async with session.get( 

42 f'{base_url}/view?{url_values}', 

43 headers={**default_headers, 'Authorization': f'Bearer {api_key}'}, 

44 ssl=AIOHTTP_CLIENT_SESSION_SSL, 

45 ) as r: 

46 r.raise_for_status() 

47 return await r.read() 

48 

49 

50def get_image_url(filename, subfolder, folder_type, base_url): 

51 log.info('get_image') 

52 data = {'filename': filename, 'subfolder': subfolder, 'type': folder_type} 

53 url_values = urllib.parse.urlencode(data) 

54 return f'{base_url}/view?{url_values}' 

55 

56 

57async def get_history(prompt_id, base_url, api_key): 

58 log.info('get_history') 

59 session = await get_session() 

60 async with session.get( 

61 f'{base_url}/history/{prompt_id}', 

62 headers={**default_headers, 'Authorization': f'Bearer {api_key}'}, 

63 ssl=AIOHTTP_CLIENT_SESSION_SSL, 

64 ) as r: 

65 r.raise_for_status() 

66 return await r.json() 

67 

68 

69async def _ws_get_images(ws, workflow, client_id, base_url, api_key): 

70 """Queue a prompt and wait on *ws* for ComfyUI to finish executing it. 

71 

72 Returns a dict of ``{'data': [{'url': ...}, ...]}``. 

73 """ 

74 prompt_id = (await queue_prompt(workflow, client_id, base_url, api_key))['prompt_id'] 

75 output_images = [] 

76 

77 async for msg in ws: 

78 if msg.type == aiohttp.WSMsgType.TEXT: 

79 message = JSONCodec.loads(msg.data) 

80 if message['type'] == 'executing': 

81 data = message['data'] 

82 if data['node'] is None and data['prompt_id'] == prompt_id: 

83 break # Execution is done 

84 elif msg.type in (aiohttp.WSMsgType.CLOSED, aiohttp.WSMsgType.ERROR): 

85 log.error(f'WebSocket closed unexpectedly: {msg.type}') 

86 break 

87 # binary messages (previews) are silently skipped 

88 

89 history = (await get_history(prompt_id, base_url, api_key))[prompt_id] 

90 for node_id in history['outputs']: 

91 node_output = history['outputs'][node_id] 

92 if node_id in workflow and workflow[node_id].get('class_type') in [ 

93 'SaveImage', 

94 'PreviewImage', 

95 ]: 

96 if 'images' in node_output: 

97 for image in node_output['images']: 

98 url = get_image_url(image['filename'], image['subfolder'], image['type'], base_url) 

99 output_images.append({'url': url}) 

100 return {'data': output_images} 

101 

102 

103async def comfyui_upload_image(image_file_item, base_url, api_key): 

104 url = f'{base_url}/api/upload/image' 

105 headers = {} 

106 

107 if api_key: 

108 headers['Authorization'] = f'Bearer {api_key}' 

109 

110 _, (filename, file_bytes, mime_type) = image_file_item 

111 

112 form = aiohttp.FormData() 

113 form.add_field('image', file_bytes, filename=filename, content_type=mime_type) 

114 form.add_field('type', 'input') # required by ComfyUI 

115 

116 session = await get_session() 

117 async with session.post(url, data=form, headers=headers, ssl=AIOHTTP_CLIENT_SESSION_SSL) as resp: 

118 resp.raise_for_status() 

119 return await resp.json() 

120 

121 

122class ComfyUINodeInput(BaseModel): 

123 type: Optional[str] = None 

124 node_ids: list[str] = [] 

125 key: Optional[str] = 'text' 

126 value: Optional[str] = None 

127 

128 

129class ComfyUIWorkflow(BaseModel): 

130 workflow: str 

131 nodes: list[ComfyUINodeInput] 

132 

133 

134class ComfyUICreateImageForm(BaseModel): 

135 workflow: ComfyUIWorkflow 

136 

137 prompt: str 

138 negative_prompt: Optional[str] = None 

139 width: int 

140 height: int 

141 n: int = 1 

142 

143 steps: Optional[int] = None 

144 seed: Optional[int] = None 

145 

146 

147def _apply_workflow_nodes(workflow, nodes, model, payload): 

148 """Mutate *workflow* dict in-place based on typed node definitions.""" 

149 for node in nodes: 

150 if node.type: 

151 if node.type == 'model': 

152 for node_id in node.node_ids: 

153 workflow[node_id]['inputs'][node.key] = model 

154 elif node.type == 'prompt': 

155 for node_id in node.node_ids: 

156 workflow[node_id]['inputs'][node.key if node.key else 'text'] = payload.prompt 

157 elif node.type == 'negative_prompt': 

158 for node_id in node.node_ids: 

159 workflow[node_id]['inputs'][node.key if node.key else 'text'] = payload.negative_prompt 

160 elif node.type == 'image': 

161 if isinstance(payload.image, list): 

162 for idx, node_id in enumerate(node.node_ids): 

163 if idx < len(payload.image): 

164 workflow[node_id]['inputs'][node.key] = payload.image[idx] 

165 else: 

166 for node_id in node.node_ids: 

167 workflow[node_id]['inputs'][node.key] = payload.image 

168 elif node.type == 'width': 

169 for node_id in node.node_ids: 

170 workflow[node_id]['inputs'][node.key if node.key else 'width'] = payload.width 

171 elif node.type == 'height': 

172 for node_id in node.node_ids: 

173 workflow[node_id]['inputs'][node.key if node.key else 'height'] = payload.height 

174 elif node.type == 'n': 

175 for node_id in node.node_ids: 

176 workflow[node_id]['inputs'][node.key if node.key else 'batch_size'] = payload.n 

177 elif node.type == 'steps': 

178 for node_id in node.node_ids: 

179 workflow[node_id]['inputs'][node.key if node.key else 'steps'] = payload.steps 

180 elif node.type == 'seed': 

181 seed = payload.seed if payload.seed else random.randint(0, 1125899906842624) 

182 for node_id in node.node_ids: 

183 workflow[node_id]['inputs'][node.key] = seed 

184 else: 

185 for node_id in node.node_ids: 

186 workflow[node_id]['inputs'][node.key] = node.value 

187 

188 

189async def comfyui_create_image(model: str, payload: ComfyUICreateImageForm, client_id, base_url, api_key): 

190 ws_url = base_url.replace('http://', 'ws://').replace('https://', 'wss://') 

191 workflow = JSONCodec.loads(payload.workflow.workflow) 

192 _apply_workflow_nodes(workflow, payload.workflow.nodes, model, payload) 

193 

194 headers = {'Authorization': f'Bearer {api_key}'} 

195 session = await get_session() 

196 

197 try: 

198 async with session.ws_connect( 

199 f'{ws_url}/ws?clientId={client_id}', 

200 headers=headers, 

201 ssl=AIOHTTP_CLIENT_SESSION_SSL, 

202 ) as ws: 

203 log.info('WebSocket connection established.') 

204 log.info('Sending workflow to WebSocket server.') 

205 log.debug('Workflow: %s', workflow) 

206 images = await _ws_get_images(ws, workflow, client_id, base_url, api_key) 

207 except aiohttp.WSServerHandshakeError as e: 

208 log.exception(f'Failed to connect to WebSocket server: {e}') 

209 return None 

210 except Exception as e: 

211 log.exception(f'Error during image generation: {e}') 

212 return None 

213 

214 return images 

215 

216 

217class ComfyUIEditImageForm(BaseModel): 

218 workflow: ComfyUIWorkflow 

219 

220 image: str | list[str] 

221 prompt: str 

222 width: Optional[int] = None 

223 height: Optional[int] = None 

224 n: Optional[int] = None 

225 

226 steps: Optional[int] = None 

227 seed: Optional[int] = None 

228 

229 

230async def comfyui_edit_image(model: str, payload: ComfyUIEditImageForm, client_id, base_url, api_key): 

231 ws_url = base_url.replace('http://', 'ws://').replace('https://', 'wss://') 

232 workflow = JSONCodec.loads(payload.workflow.workflow) 

233 _apply_workflow_nodes(workflow, payload.workflow.nodes, model, payload) 

234 

235 headers = {'Authorization': f'Bearer {api_key}'} 

236 session = await get_session() 

237 

238 try: 

239 async with session.ws_connect( 

240 f'{ws_url}/ws?clientId={client_id}', 

241 headers=headers, 

242 ssl=AIOHTTP_CLIENT_SESSION_SSL, 

243 ) as ws: 

244 log.info('WebSocket connection established.') 

245 log.info('Sending workflow to WebSocket server.') 

246 log.debug('Workflow: %s', workflow) 

247 images = await _ws_get_images(ws, workflow, client_id, base_url, api_key) 

248 except aiohttp.WSServerHandshakeError as e: 

249 log.exception(f'Failed to connect to WebSocket server: {e}') 

250 return None 

251 except Exception as e: 

252 log.exception(f'Error during image editing: {e}') 

253 return None 

254 

255 return images