Coverage for open_webui/utils/filter.py: 12%
146 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 inspect
2import logging
4from open_webui.env import ENABLE_PLUGINS
5from open_webui.models.functions import Functions
6from open_webui.utils.plugin import get_function_module_from_cache
8log = logging.getLogger(__name__)
11class FilterContext:
12 def __init__(self):
13 self.active_filters = None
14 self.valves_by_id = None
15 self.function_valves = {}
16 self.user_valves = {}
18 async def get_active_filters(self):
19 if self.active_filters is None:
20 self.active_filters = await Functions.get_active_filter_ids()
21 return self.active_filters
23 async def get_function_valves(self, filter_ids, filter_id, Valves):
24 if filter_id not in self.function_valves:
25 if self.valves_by_id is None:
26 self.valves_by_id = await Functions.get_function_valves_by_ids(filter_ids)
27 valves = self.valves_by_id.get(filter_id)
28 self.function_valves[filter_id] = Valves(**(valves if valves else {}))
29 return self.function_valves[filter_id]
31 async def get_user_valves(self, filter_id, user_id, UserValves):
32 user_valves_key = (filter_id, user_id)
33 if user_valves_key not in self.user_valves:
34 self.user_valves[user_valves_key] = await get_user_valves(filter_id, user_id, UserValves)
35 return self.user_valves[user_valves_key]
38def get_filter_context(request):
39 if not hasattr(request.state, 'filter_context'):
40 request.state.filter_context = FilterContext()
41 return request.state.filter_context
44async def get_user_valves(filter_id, user_id, UserValves):
45 user_valves_data = await Functions.get_user_valves_by_id_and_user_id(filter_id, user_id)
46 return UserValves(**(user_valves_data if user_valves_data else {}))
49async def get_function_module(request, function_id, load_from_db=True, function=None):
50 """
51 Get the function module by its ID.
52 """
53 function_module, _, _ = await get_function_module_from_cache(
54 request, function_id, function=function, load_from_db=load_from_db
55 )
56 return function_module
59def get_model_filter_ids(model, active_filters):
60 filter_ids = [fid for fid, is_global in active_filters if is_global]
61 if isinstance(model, dict) and 'info' in model and 'meta' in model['info']:
62 filter_ids.extend(model['info']['meta'].get('filterIds', []))
63 filter_ids = list(set(filter_ids))
64 active_filter_ids = {fid for fid, _ in active_filters}
65 return [fid for fid in filter_ids if fid in active_filter_ids]
68async def resolve_filter_pipeline(request, model: dict, enabled_filter_ids: list = None):
69 if not ENABLE_PLUGINS:
70 return [], []
72 active_filters = await get_filter_context(request).get_active_filters()
73 filter_ids = get_model_filter_ids(model, active_filters)
74 functions_by_id = {function.id: function for function in await Functions.get_functions_by_ids(filter_ids)}
76 async def get_active_status(filter_id):
77 function_module = await get_function_module(request, filter_id, function=functions_by_id.get(filter_id))
79 if getattr(function_module, 'toggle', None):
80 return filter_id in (enabled_filter_ids or set())
82 return True
84 # Pre-compute active status for each filter (async functions can't be used in set comprehensions)
85 filter_ids = [fid for fid in filter_ids if await get_active_status(fid)]
86 valves_by_id = await Functions.get_function_valves_by_ids(filter_ids)
88 async def get_priority(function_id):
89 try:
90 function_module = await get_function_module(request, function_id, function=functions_by_id.get(function_id))
91 if function_module and hasattr(function_module, 'Valves'):
92 valves_db = valves_by_id.get(function_id)
93 valves = function_module.Valves(**(valves_db if valves_db else {}))
94 return getattr(valves, 'priority', 0)
95 except Exception:
96 pass
97 return 0
99 # Pre-compute priorities (async functions can't be used in sort keys)
100 priorities = {}
101 for fid in filter_ids:
102 priorities[fid] = await get_priority(fid)
103 filter_ids.sort(key=lambda fid: (priorities.get(fid, 0), fid))
105 filter_functions = [functions_by_id[fid] for fid in filter_ids if fid in functions_by_id]
106 return filter_ids, filter_functions
109async def get_sorted_filter_ids(request, model: dict, enabled_filter_ids: list = None):
110 filter_ids, _ = await resolve_filter_pipeline(request, model, enabled_filter_ids)
112 return filter_ids
115async def get_filter_functions(request, model: dict, enabled_filter_ids: list = None):
116 _, filter_functions = await resolve_filter_pipeline(request, model, enabled_filter_ids)
117 return filter_functions
120async def apply_filter_valves(function_module, filter_context, valves_by_id, filter_ids, filter_id):
121 if not (hasattr(function_module, 'valves') and hasattr(function_module, 'Valves')):
122 return valves_by_id
124 if filter_context is not None:
125 function_module.valves = await filter_context.get_function_valves(filter_ids, filter_id, function_module.Valves)
126 return valves_by_id
128 if valves_by_id is None:
129 valves_by_id = await Functions.get_function_valves_by_ids(filter_ids)
130 valves = valves_by_id.get(filter_id)
131 function_module.valves = function_module.Valves(**(valves if valves else {}))
132 return valves_by_id
135def get_filter_params(sig, filter_id, filter_type, form_data, extra_params):
136 params = {'event': form_data} if filter_type == 'stream' else {'body': form_data}
137 return params | {
138 k: v
139 for k, v in {
140 **extra_params,
141 '__id__': filter_id,
142 }.items()
143 if k in sig.parameters
144 }
147async def apply_user_valves(function_module, filter_context, filter_id, params):
148 if '__user__' not in params or not hasattr(function_module, 'UserValves'):
149 return
151 user_id = params['__user__'].get('id')
152 if filter_context is not None:
153 user_valves = await filter_context.get_user_valves(filter_id, user_id, function_module.UserValves)
154 else:
155 user_valves = await get_user_valves(filter_id, user_id, function_module.UserValves)
156 params['__user__']['valves'] = user_valves
159async def run_filter_handler(handler, params):
160 if inspect.iscoroutinefunction(handler):
161 return await handler(**params)
162 return handler(**params)
165async def process_filter_function(
166 request,
167 function,
168 filter_type,
169 form_data,
170 extra_params,
171 filter_context,
172 valves_by_id,
173 filter_ids,
174):
175 filter_id = function.id
177 function_module = await get_function_module(
178 request, filter_id, load_from_db=(filter_type != 'stream'), function=function
179 )
180 handler = getattr(function_module, filter_type, None)
181 if not callable(handler):
182 return form_data, valves_by_id, None
184 skip_files = (
185 function_module.file_handler if filter_type == 'inlet' and hasattr(function_module, 'file_handler') else None
186 )
188 try:
189 valves_by_id = await apply_filter_valves(function_module, filter_context, valves_by_id, filter_ids, filter_id)
190 sig = inspect.signature(handler)
191 params = get_filter_params(sig, filter_id, filter_type, form_data, extra_params)
193 if '__user__' in sig.parameters:
194 try:
195 await apply_user_valves(function_module, filter_context, filter_id, params)
196 except Exception:
197 log.exception('Failed to get user valves for filter %s', filter_id)
199 form_data = await run_filter_handler(handler, params)
200 except Exception:
201 if filter_type == 'inlet':
202 log.debug('Error in inlet filter %s', filter_id, exc_info=True)
203 else:
204 log.exception('Error in %s filter %s', filter_type, filter_id)
205 raise
207 return form_data, valves_by_id, skip_files
210# Grant these filters the discernment to pass what serves
211# and refuse what harms, for every soul in the house.
212async def process_filter_functions(
213 request,
214 filter_context,
215 filter_functions,
216 filter_type,
217 form_data,
218 extra_params,
219):
220 if not ENABLE_PLUGINS:
221 return form_data, {}
223 skip_files = None
224 valves_by_id = None
225 filter_ids = [function.id for function in filter_functions if function]
227 for function in filter_functions:
228 if not function:
229 continue
231 form_data, valves_by_id, file_handler = await process_filter_function(
232 request,
233 function,
234 filter_type,
235 form_data,
236 extra_params,
237 filter_context,
238 valves_by_id,
239 filter_ids,
240 )
241 skip_files = skip_files or file_handler
243 # Handle file cleanup for inlet
244 if skip_files:
245 if 'files' in form_data.get('metadata', {}):
246 del form_data['metadata']['files']
247 if 'files' in form_data:
248 del form_data['files']
250 return form_data, {}