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
« 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
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
43logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
44log = logging.getLogger(__name__)
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')
57 metadata = form_data.pop('metadata', {})
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
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 )
70 channel = f'{user_id}:{session_id}:{request_id}'
71 logging.info('WebSocket channel: %s', channel)
73 if form_data.get('stream'):
74 queue = asyncio.Queue()
75 EVENT_QUEUES[channel] = queue
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 )
91 log.info('res: %s', res)
93 status = res.get('status', False)
94 except BaseException:
95 EVENT_QUEUES.pop(channel, None)
96 raise
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
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)
120 # Define a background task to run the event generator
121 async def background():
122 EVENT_QUEUES.pop(channel, None)
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 )
142 if 'error' in res and res['error']:
143 raise Exception(res['error'])
145 return res
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
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
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 }
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
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')
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
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)
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 ]
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)
237 form_data['model'] = selected_model_id
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)
245 if selected_model_id:
246 if form_data.get('stream') == True:
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
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}
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 )
308chat_completion = generate_chat_completion
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)
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
323 data = form_data
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')
328 model_id = data['model']
329 if model_id not in models:
330 raise Exception('Model not found')
332 model = models[model_id]
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}')
341 if not data.get('id'):
342 raise Exception('Missing message id')
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 }
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 }
361 try:
362 filter_functions = await get_filter_functions(request, model, metadata.get('filter_ids', []))
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}')