Coverage for open_webui/utils/chat.py: 16%

169 statements  

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

1import asyncio 

2import logging 

3import random 

4import sys 

5import time 

6import uuid 

7from typing import Any, Optional 

8 

9from aiocache import cached 

10from fastapi import HTTPException, Request, status 

11from open_webui.env import BYPASS_MODEL_ACCESS_CONTROL, GLOBAL_LOG_LEVEL 

12from open_webui.functions import generate_function_chat_completion 

13from open_webui.models.models import Models 

14from open_webui.models.users import UserModel 

15from open_webui.routers.ollama import ( 

16 generate_chat_completion as generate_ollama_chat_completion, 

17) 

18from open_webui.routers.openai import ( 

19 generate_chat_completion as generate_openai_chat_completion, 

20) 

21from open_webui.routers.pipelines import ( 

22 process_pipeline_inlet_filter, 

23 process_pipeline_outlet_filter, 

24) 

25from open_webui.socket.main import ( 

26 EVENT_QUEUES, 

27 get_event_call, 

28 get_event_emitter, 

29) 

30from open_webui.utils.filter import ( 

31 get_filter_functions, 

32 process_filter_functions, 

33) 

34from open_webui.utils.json_codec import JSONCodec 

35from open_webui.utils.models import check_model_access, get_all_models 

36from open_webui.utils.payload import convert_payload_openai_to_ollama 

37from open_webui.utils.response import ( 

38 convert_response_ollama_to_openai, 

39 convert_streaming_response_ollama_to_openai, 

40) 

41from starlette.responses import JSONResponse, Response, StreamingResponse 

42 

43logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) 

44log = logging.getLogger(__name__) 

45 

46 

47# When the question has been asked, let silence not be the 

48# answer. But if the answer must wait, let it come honest. 

49async def generate_direct_chat_completion( 

50 request: Request, 

51 form_data: dict, 

52 user: Any, 

53 models: dict, 

54): 

55 log.info('generate_direct_chat_completion') 

56 

57 metadata = form_data.pop('metadata', {}) 

58 

59 user_id = metadata.get('user_id') 

60 session_id = metadata.get('session_id') 

61 request_id = str(uuid.uuid4()) # Generate a unique request ID 

62 

63 event_caller = await get_event_call(metadata) 

64 if event_caller is None: 

65 raise Exception( 

66 'Direct connection requires an active WebSocket session; ' 

67 'cannot generate completion in this context (e.g. background task).' 

68 ) 

69 

70 channel = f'{user_id}:{session_id}:{request_id}' 

71 logging.info('WebSocket channel: %s', channel) 

72 

73 if form_data.get('stream'): 

74 queue = asyncio.Queue() 

75 EVENT_QUEUES[channel] = queue 

76 

77 # Start processing chat completion in background 

78 try: 

79 res = await event_caller( 

80 { 

81 'type': 'request:chat:completion', 

82 'data': { 

83 'form_data': form_data, 

84 'model': models[form_data['model']], 

85 'channel': channel, 

86 'session_id': session_id, 

87 }, 

88 } 

89 ) 

90 

91 log.info('res: %s', res) 

92 

93 status = res.get('status', False) 

94 except BaseException: 

95 EVENT_QUEUES.pop(channel, None) 

96 raise 

97 

98 if status: 

99 # Define a generator to stream responses 

100 async def event_generator(): 

101 try: 

102 while True: 

103 data = await queue.get() # Wait for new messages 

104 if isinstance(data, dict): 

105 if 'done' in data and data['done']: 

106 break # Stop streaming when 'done' is received 

107 

108 yield f'data: {JSONCodec.dumps(data)}\n\n' 

109 elif isinstance(data, str): 

110 if 'data:' in data: 

111 yield f'{data}\n\n' 

112 else: 

113 yield f'data: {data}\n\n' 

114 except Exception as e: 

115 log.debug('Error in event generator: %s', e) 

116 pass 

117 finally: 

118 EVENT_QUEUES.pop(channel, None) 

119 

120 # Define a background task to run the event generator 

121 async def background(): 

122 EVENT_QUEUES.pop(channel, None) 

123 

124 # Return the streaming response 

125 return StreamingResponse(event_generator(), media_type='text/event-stream', background=background) 

126 else: 

127 EVENT_QUEUES.pop(channel, None) 

128 raise Exception(str(res)) 

129 else: 

130 res = await event_caller( 

131 { 

132 'type': 'request:chat:completion', 

133 'data': { 

134 'form_data': form_data, 

135 'model': models[form_data['model']], 

136 'channel': channel, 

137 'session_id': session_id, 

138 }, 

139 } 

140 ) 

141 

142 if 'error' in res and res['error']: 

143 raise Exception(res['error']) 

144 

145 return res 

146 

147 

148async def generate_chat_completion( 

149 request: Request, 

150 form_data: dict, 

151 user: Any, 

152 bypass_filter: bool = False, 

153 bypass_system_prompt: bool = False, 

154): 

155 log.debug('generate_chat_completion: %s', form_data) 

156 if BYPASS_MODEL_ACCESS_CONTROL: 

157 bypass_filter = True 

158 

159 # Propagate bypass_filter and bypass_system_prompt via request.state so that 

160 # downstream route handlers (openai/ollama) can read them without exposing 

161 # them as query parameters. 

162 request.state.bypass_filter = bypass_filter 

163 request.state.bypass_system_prompt = bypass_system_prompt 

164 

165 if hasattr(request.state, 'metadata'): 

166 if 'metadata' not in form_data: 

167 form_data['metadata'] = request.state.metadata 

168 else: 

169 form_data['metadata'] = { 

170 **form_data['metadata'], 

171 **request.state.metadata, 

172 } 

173 

174 if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): 

175 # Merge the direct connection model into server models so that 

176 # task functions (title, tags, etc.) can resolve a server-side 

177 # task model while still having the direct model available. 

178 # dict(...items()) is one HGETALL on a Redis-backed pool; ``{**pool}`` 

179 # would issue HKEYS plus one HGET per model. 

180 models = { 

181 **dict(request.app.state.MODELS.items()), 

182 request.state.model['id']: request.state.model, 

183 } 

184 log.debug('direct connection to model: %s', request.state.model['id']) 

185 else: 

186 models = request.app.state.MODELS 

187 

188 model_id = form_data['model'] 

189 # Single lookup — membership check plus getitem would be two Redis 

190 # round trips on a Redis-backed model pool. 

191 model = models.get(model_id) 

192 if model is None: 

193 raise Exception('Model not found') 

194 

195 if getattr(request.state, 'direct', False) and model_id == getattr(request.state, 'model', {}).get('id'): 

196 return await generate_direct_chat_completion(request, form_data, user=user, models=models) 

197 else: 

198 # Check if user has access to the model 

199 if not bypass_filter and user.role == 'user': 

200 try: 

201 await check_model_access(user, model) 

202 except Exception as e: 

203 raise e 

204 

205 # Arena model — sub-model was already resolved by process_chat_payload. 

206 # Inject selected_model_id into the response for the frontend. 

207 metadata = form_data.get('metadata', {}) 

208 selected_model_id = metadata.pop('selected_model_id', None) 

209 # Also clear from request.state.metadata to prevent the merge at 

210 # lines 177-179 from re-adding it on the recursive call. 

211 if hasattr(request.state, 'metadata'): 

212 request.state.metadata.pop('selected_model_id', None) 

213 

214 # Fallback: if generate_chat_completion is called with an arena model 

215 # from a path that did NOT go through process_chat_payload (e.g., 

216 # background tasks for title/follow-up/tags generation), resolve now. 

217 if not selected_model_id and model.get('owned_by') == 'arena': 

218 model_ids = model.get('info', {}).get('meta', {}).get('model_ids') 

219 filter_mode = model.get('info', {}).get('meta', {}).get('filter_mode') 

220 if model_ids and filter_mode == 'exclude': 

221 model_ids = [ 

222 available_model['id'] 

223 for available_model in list(request.app.state.MODELS.values()) 

224 if available_model.get('owned_by') != 'arena' and available_model['id'] not in model_ids 

225 ] 

226 

227 if isinstance(model_ids, list) and model_ids: 

228 selected_model_id = random.choice(model_ids) 

229 else: 

230 model_ids = [ 

231 available_model['id'] 

232 for available_model in list(request.app.state.MODELS.values()) 

233 if available_model.get('owned_by') != 'arena' 

234 ] 

235 selected_model_id = random.choice(model_ids) 

236 

237 form_data['model'] = selected_model_id 

238 

239 # bypass_filter recursion below skips the line-200 check; gate the resolved model here. 

240 if not bypass_filter and user.role == 'user': 

241 selected_model = request.app.state.MODELS.get(selected_model_id) 

242 if selected_model: 

243 await check_model_access(user, selected_model) 

244 

245 if selected_model_id: 

246 if form_data.get('stream') == True: 

247 

248 async def stream_wrapper(stream): 

249 yield f'data: {JSONCodec.dumps({"selected_model_id": selected_model_id})}\n\n' 

250 async for chunk in stream: 

251 yield chunk 

252 

253 response = await generate_chat_completion( 

254 request, 

255 form_data, 

256 user, 

257 bypass_filter=True, 

258 bypass_system_prompt=bypass_system_prompt, 

259 ) 

260 # Upstream errors come back as a response object. 

261 if not isinstance(response, StreamingResponse): 

262 return response 

263 return StreamingResponse( 

264 stream_wrapper(response.body_iterator), 

265 media_type='text/event-stream', 

266 background=response.background, 

267 ) 

268 else: 

269 response = await generate_chat_completion( 

270 request, 

271 form_data, 

272 user, 

273 bypass_filter=True, 

274 bypass_system_prompt=bypass_system_prompt, 

275 ) 

276 if not isinstance(response, dict): 

277 return response 

278 return {**response, 'selected_model_id': selected_model_id} 

279 

280 if model.get('pipe'): 

281 # Below does not require bypass_filter because this is the only route the uses this function and it is already bypassing the filter 

282 return await generate_function_chat_completion(request, form_data, user=user, models=models) 

283 if model.get('owned_by') == 'ollama': 

284 # Using /ollama/api/chat endpoint 

285 form_data = convert_payload_openai_to_ollama(form_data) 

286 response = await generate_ollama_chat_completion( 

287 request=request, 

288 form_data=form_data, 

289 user=user, 

290 ) 

291 if form_data.get('stream'): 

292 response.headers['content-type'] = 'text/event-stream' 

293 return StreamingResponse( 

294 convert_streaming_response_ollama_to_openai(response), 

295 headers=dict(response.headers), 

296 background=response.background, 

297 ) 

298 else: 

299 return convert_response_ollama_to_openai(response) 

300 else: 

301 return await generate_openai_chat_completion( 

302 request=request, 

303 form_data=form_data, 

304 user=user, 

305 ) 

306 

307 

308chat_completion = generate_chat_completion 

309 

310 

311async def chat_completed(request: Request, form_data: dict, user: Any): 

312 if not request.app.state.MODELS: 312 ↛ 315line 312 didn't jump to line 315 because the condition on line 312 was always true

313 await get_all_models(request, user=user) 

314 

315 if getattr(request.state, 'direct', False) and hasattr(request.state, 'model'): 315 ↛ 316line 315 didn't jump to line 316 because the condition on line 315 was never true

316 models = { 

317 **dict(request.app.state.MODELS.items()), 

318 request.state.model['id']: request.state.model, 

319 } 

320 else: 

321 models = request.app.state.MODELS 

322 

323 data = form_data 

324 

325 if not data.get('id'): 325 ↛ 328line 325 didn't jump to line 328 because the condition on line 325 was always true

326 raise Exception('Missing message id') 

327 

328 model_id = data['model'] 

329 if model_id not in models: 

330 raise Exception('Model not found') 

331 

332 model = models[model_id] 

333 

334 try: 

335 data = await process_pipeline_outlet_filter(request, data, user, models) 

336 except HTTPException: 

337 raise 

338 except Exception as e: 

339 raise Exception(f'Error: {e}') 

340 

341 if not data.get('id'): 

342 raise Exception('Missing message id') 

343 

344 metadata = { 

345 'chat_id': data['chat_id'], 

346 'message_id': data['id'], 

347 'filter_ids': data.get('filter_ids', []), 

348 'session_id': data['session_id'], 

349 'user_id': user.id, 

350 } 

351 

352 extra_params = { 

353 '__event_emitter__': await get_event_emitter(metadata), 

354 '__event_call__': await get_event_call(metadata), 

355 '__user__': user.model_dump() if isinstance(user, UserModel) else {}, 

356 '__metadata__': metadata, 

357 '__request__': request, 

358 '__model__': model, 

359 } 

360 

361 try: 

362 filter_functions = await get_filter_functions(request, model, metadata.get('filter_ids', [])) 

363 

364 result, _ = await process_filter_functions( 

365 request=request, 

366 filter_context=None, 

367 filter_functions=filter_functions, 

368 filter_type='outlet', 

369 form_data=data, 

370 extra_params=extra_params, 

371 ) 

372 return result 

373 except Exception as e: 

374 raise Exception(f'Error: {e}')