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

1import inspect 

2import logging 

3 

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 

7 

8log = logging.getLogger(__name__) 

9 

10 

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 = {} 

17 

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 

22 

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] 

30 

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] 

36 

37 

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 

42 

43 

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 {})) 

47 

48 

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 

57 

58 

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] 

66 

67 

68async def resolve_filter_pipeline(request, model: dict, enabled_filter_ids: list = None): 

69 if not ENABLE_PLUGINS: 

70 return [], [] 

71 

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)} 

75 

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)) 

78 

79 if getattr(function_module, 'toggle', None): 

80 return filter_id in (enabled_filter_ids or set()) 

81 

82 return True 

83 

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) 

87 

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 

98 

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)) 

104 

105 filter_functions = [functions_by_id[fid] for fid in filter_ids if fid in functions_by_id] 

106 return filter_ids, filter_functions 

107 

108 

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) 

111 

112 return filter_ids 

113 

114 

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 

118 

119 

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 

123 

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 

127 

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 

133 

134 

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 } 

145 

146 

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 

150 

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 

157 

158 

159async def run_filter_handler(handler, params): 

160 if inspect.iscoroutinefunction(handler): 

161 return await handler(**params) 

162 return handler(**params) 

163 

164 

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 

176 

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 

183 

184 skip_files = ( 

185 function_module.file_handler if filter_type == 'inlet' and hasattr(function_module, 'file_handler') else None 

186 ) 

187 

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) 

192 

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) 

198 

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 

206 

207 return form_data, valves_by_id, skip_files 

208 

209 

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, {} 

222 

223 skip_files = None 

224 valves_by_id = None 

225 filter_ids = [function.id for function in filter_functions if function] 

226 

227 for function in filter_functions: 

228 if not function: 

229 continue 

230 

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 

242 

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'] 

249 

250 return form_data, {}