Coverage for open_webui/models/config.py: 43%

240 statements  

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

1"""Database-backed configuration with per-key storage. 

2 

3Replaces the old single-row JSON blob machinery with a simple per-key model 

4mirroring cptr's Config. 

5 

6Each config key is stored as its own row: key TEXT PK, value JSON. 

7Reads are direct DB lookups. Writes are explicit awaited upserts that raise on 

8failure (no more fire-and-forget create_task). 

9""" 

10 

11from __future__ import annotations 

12 

13import logging 

14import time 

15from typing import Any, ClassVar 

16 

17from fastapi.encoders import jsonable_encoder 

18from open_webui.internal.db import Base, get_async_db 

19from sqlalchemy import JSON, BigInteger, Column, Text, delete, select 

20 

21log = logging.getLogger(__name__) 

22 

23API_CONFIG_KEYS = ('openai.api_configs', 'ollama.api_configs') 

24DICT_CONFIG_KEY_ALIASES = { 

25 'openai.api_configs': ('OPENAI_API_CONFIGS',), 

26 'ollama.api_configs': ('OLLAMA_API_CONFIGS',), 

27 'rag.mineru_params': ('MINERU_PARAMS',), 

28 'rag.docling_params': ('DOCLING_PARAMS',), 

29 'web.search.linkup_search_params': ('LINKUP_SEARCH_PARAMS',), 

30 'image_generation.automatic1111.api_params': ('AUTOMATIC1111_PARAMS',), 

31 'image_generation.openai.params': ('IMAGES_OPENAI_API_PARAMS',), 

32 'audio.tts.openai.params': ('AUDIO_TTS_OPENAI_PARAMS',), 

33 'models.default_metadata': ('DEFAULT_MODEL_METADATA',), 

34 'models.default_params': ('DEFAULT_MODEL_PARAMS',), 

35 'task.model.params': ('TASK_MODEL_PARAMS',), 

36 'ui.default_interface_settings': ('DEFAULT_INTERFACE_SETTINGS',), 

37 'user.permissions': ('USER_PERMISSIONS',), 

38} 

39DICT_CONFIG_KEYS = tuple(DICT_CONFIG_KEY_ALIASES) 

40API_CONFIG_FIELDS = ( 

41 'enable', 

42 'key', 

43 'prefix_id', 

44 'tags', 

45 'model_ids', 

46 'connection_type', 

47 'provider', 

48 'auth_type', 

49 'headers', 

50 'azure', 

51 'api_type', 

52 'api_version', 

53 'extra_params', 

54 'passthrough_params', 

55) 

56 

57 

58def _split_api_config_fragment(fragment: str) -> tuple[str, list[str]] | None: 

59 if not fragment: 

60 return None 

61 

62 first, _, rest = fragment.partition('.') 

63 if first.isdigit() and rest: 

64 return first, rest.split('.') 

65 

66 match: tuple[int, str] | None = None 

67 for field in API_CONFIG_FIELDS: 

68 marker = f'.{field}' 

69 marker_index = fragment.rfind(marker) 

70 if marker_index != -1 and (match is None or marker_index > match[0]): 

71 match = (marker_index, field) 

72 

73 if match: 

74 marker_index, field = match 

75 connection_key = fragment[:marker_index] 

76 field_path = fragment[marker_index + 1 :] 

77 if connection_key: 

78 return connection_key, field_path.split('.') 

79 

80 return None 

81 

82 

83def _assign_path(target: dict, path: list[str], value: Any) -> None: 

84 current = target 

85 for part in path[:-1]: 

86 next_value = current.get(part) 

87 if not isinstance(next_value, dict): 

88 next_value = {} 

89 current[part] = next_value 

90 current = next_value 

91 current[path[-1]] = value 

92 

93 

94def _json_value(value: Any) -> Any: 

95 return jsonable_encoder(value) 

96 

97 

98# ── Model ──────────────────────────────────────────────────────────────────── 

99 

100 

101class Config(Base): 

102 """Per-key config storage. Each row is one config key.""" 

103 

104 __tablename__ = 'config' 

105 

106 key = Column(Text, primary_key=True) 

107 value = Column(JSON, nullable=False) 

108 updated_at = Column(BigInteger, nullable=True) 

109 

110 DEFAULTS: ClassVar[dict[str, Any]] = {} 

111 PERSISTENT_ENABLED: ClassVar[bool] = True 

112 OAUTH_PERSISTENT_ENABLED: ClassVar[bool] = False 

113 

114 # ── Class methods ──────────────────────────────────────── 

115 

116 @classmethod 

117 def configure( 

118 cls, 

119 *, 

120 defaults: dict[str, Any] | None = None, 

121 enable_persistent: bool = True, 

122 enable_oauth_persistent: bool = False, 

123 ) -> None: 

124 cls.DEFAULTS = dict(defaults or {}) 

125 cls.PERSISTENT_ENABLED = enable_persistent 

126 cls.OAUTH_PERSISTENT_ENABLED = enable_oauth_persistent 

127 

128 @classmethod 

129 def default_value(cls, key: str, default: Any = None) -> Any: 

130 return cls.DEFAULTS.get(key, default) 

131 

132 @classmethod 

133 def persistent_enabled_for(cls, key: str) -> bool: 

134 if not cls.PERSISTENT_ENABLED: 134 ↛ 135line 134 didn't jump to line 135 because the condition on line 134 was never true

135 return False 

136 if key.startswith('oauth.') and not cls.OAUTH_PERSISTENT_ENABLED: 

137 return False 

138 return True 

139 

140 @staticmethod 

141 async def get(key: str, default: Any = None) -> Any: 

142 """Get a config value by key. Returns default if not set.""" 

143 if not Config.persistent_enabled_for(key): 143 ↛ 144line 143 didn't jump to line 144 because the condition on line 143 was never true

144 return Config.default_value(key, default) 

145 async with get_async_db() as db: 

146 row = await db.get(Config, key) 

147 return row.value if row else Config.default_value(key, default) 

148 

149 @staticmethod 

150 async def get_many(*keys: str) -> dict: 

151 """Get multiple config values. Returns {key: value} for keys that exist.""" 

152 disabled_values = { 

153 key: Config.default_value(key) 

154 for key in keys 

155 if not Config.persistent_enabled_for(key) and key in Config.DEFAULTS 

156 } 

157 enabled_keys = {key for key in keys if Config.persistent_enabled_for(key)} 

158 if not enabled_keys: 

159 return disabled_values 

160 async with get_async_db() as db: 

161 result = await db.execute(select(Config).where(Config.key.in_(enabled_keys))) 

162 values = {row.key: row.value for row in result.scalars().all()} 

163 return { 

164 key: values.get(key, Config.default_value(key)) 

165 for key in keys 

166 if key in values or key in Config.DEFAULTS or key in disabled_values 

167 } 

168 

169 @staticmethod 

170 async def get_namespace(namespace: str) -> dict: 

171 """Get all config keys under a dotted namespace.""" 

172 default_values = { 

173 key: value 

174 for key, value in Config.DEFAULTS.items() 

175 if key.startswith(f'{namespace}.') and not Config.persistent_enabled_for(key) 

176 } 

177 if not Config.PERSISTENT_ENABLED: 177 ↛ 178line 177 didn't jump to line 178 because the condition on line 177 was never true

178 return default_values 

179 async with get_async_db() as db: 

180 result = await db.execute(select(Config).where(Config.key.like(f'{namespace}.%'))) 

181 values = {row.key: row.value for row in result.scalars().all()} 

182 values.update(default_values) 

183 return values 

184 

185 @staticmethod 

186 async def get_all() -> dict: 

187 """Get all config as {key: value}.""" 

188 if not Config.PERSISTENT_ENABLED: 188 ↛ 189line 188 didn't jump to line 189 because the condition on line 188 was never true

189 return dict(Config.DEFAULTS) 

190 async with get_async_db() as db: 

191 result = await db.execute(select(Config)) 

192 values = {row.key: row.value for row in result.scalars().all()} 

193 if not Config.OAUTH_PERSISTENT_ENABLED: 

194 values.update({key: value for key, value in Config.DEFAULTS.items() if key.startswith('oauth.')}) 

195 return values 

196 

197 @staticmethod 

198 async def upsert(updates: dict) -> None: 

199 """Upsert multiple config key-value pairs. Raises on failure.""" 

200 persistent_updates = {} 

201 for key, value in updates.items(): 

202 value = _json_value(value) 

203 if Config.persistent_enabled_for(key): 

204 persistent_updates[key] = value 

205 else: 

206 Config.DEFAULTS[key] = value 

207 

208 if not persistent_updates: 

209 return 

210 

211 async with get_async_db() as db: 

212 now = int(time.time()) 

213 for key, value in persistent_updates.items(): 213 ↛ 220line 213 didn't jump to line 220 because the loop on line 213 didn't complete

214 existing = await db.get(Config, key) 

215 if existing: 

216 existing.value = value 

217 existing.updated_at = now 

218 else: 

219 db.add(Config(key=key, value=value, updated_at=now)) 

220 await db.commit() 

221 

222 @staticmethod 

223 async def delete(key: str) -> bool: 

224 """Delete a config key. Returns True if it existed.""" 

225 async with get_async_db() as db: 

226 row = await db.get(Config, key) 

227 if row: 

228 await db.delete(row) 

229 await db.commit() 

230 return True 

231 return False 

232 

233 @staticmethod 

234 async def clear() -> None: 

235 """Delete all config rows.""" 

236 async with get_async_db() as db: 

237 await db.execute(delete(Config)) 

238 await db.commit() 

239 

240 @staticmethod 

241 async def seed_defaults(defaults: dict) -> None: 

242 """Insert keys that don't yet exist in the DB. 

243 

244 Called at startup to ensure all known config keys have values. 

245 Existing DB values take precedence over defaults. 

246 """ 

247 async with get_async_db() as db: 

248 result = await db.execute(select(Config.key)) 

249 existing_keys = {row[0] for row in result.all()} 

250 

251 now = int(time.time()) 

252 new_count = 0 

253 for key, value in defaults.items(): 

254 # Skip keys the DB is not authoritative for (e.g. oauth.* while 

255 # ENABLE_OAUTH_PERSISTENT_CONFIG is off), matching the read paths. 

256 if not Config.persistent_enabled_for(key): 

257 continue 

258 if key not in existing_keys: 

259 value = _json_value(value) 

260 db.add(Config(key=key, value=value, updated_at=now)) 

261 existing_keys.add(key) 

262 new_count += 1 

263 

264 if new_count: 

265 await db.commit() 

266 log.info('Seeded %d new config defaults', new_count) 

267 

268 @staticmethod 

269 async def rename_prefix(old_prefix: str, new_prefix: str) -> None: 

270 """Move persisted config keys from one dotted prefix to another.""" 

271 if not Config.PERSISTENT_ENABLED: 271 ↛ 272line 271 didn't jump to line 272 because the condition on line 271 was never true

272 return 

273 

274 async with get_async_db() as db: 

275 result = await db.execute(select(Config).where(Config.key.like(f'{old_prefix}.%'))) 

276 rows = result.scalars().all() 

277 if not rows: 

278 return 

279 

280 now = int(time.time()) 

281 moved_count = 0 

282 deleted_count = 0 

283 for row in rows: 

284 new_key = f'{new_prefix}.{row.key.removeprefix(f"{old_prefix}.")}' 

285 existing = await db.get(Config, new_key) 

286 if existing is None: 

287 db.add(Config(key=new_key, value=row.value, updated_at=now)) 

288 moved_count += 1 

289 else: 

290 deleted_count += 1 

291 await db.delete(row) 

292 

293 await db.commit() 

294 log.info( 

295 'Renamed %d config keys from %s.* to %s.*; deleted %d old duplicates', 

296 moved_count, 

297 old_prefix, 

298 new_prefix, 

299 deleted_count, 

300 ) 

301 

302 @staticmethod 

303 async def repair_config_rows() -> None: 

304 """Repair known legacy config row shapes.""" 

305 if not Config.PERSISTENT_ENABLED: 305 ↛ 306line 305 didn't jump to line 306 because the condition on line 305 was never true

306 return 

307 

308 async with get_async_db() as db: 

309 repaired_keys: list[str] = [] 

310 orphan_keys: list[str] = [] 

311 default_model_keys: list[str] = [] 

312 now = int(time.time()) 

313 

314 for config_key, aliases in DICT_CONFIG_KEY_ALIASES.items(): 314 ↛ 365line 314 didn't jump to line 365 because the loop on line 314 didn't complete

315 prefixes = (config_key, *aliases) 

316 rows = [] 

317 for key_prefix in prefixes: 317 ↛ 320line 317 didn't jump to line 320 because the loop on line 317 didn't complete

318 result = await db.execute(select(Config).where(Config.key.like(f'{key_prefix}.%'))) 

319 rows.extend(result.scalars().all()) 

320 if not rows: 320 ↛ 323line 320 didn't jump to line 323 because the condition on line 320 was always true

321 continue 

322 

323 existing = await db.get(Config, config_key) 

324 repaired = existing.value if existing and isinstance(existing.value, dict) else {} 

325 

326 repaired_any = False 

327 for row in rows: 327 ↛ 355line 327 didn't jump to line 355 because the loop on line 327 didn't complete

328 fragment = None 

329 for key_prefix in prefixes: 

330 prefix = f'{key_prefix}.' 

331 if row.key.startswith(prefix): 

332 fragment = row.key.removeprefix(prefix) 

333 break 

334 if fragment is None: 

335 continue 

336 

337 if config_key in API_CONFIG_KEYS: 

338 split = _split_api_config_fragment(fragment) 

339 if not split: 

340 continue 

341 object_key, field_path = split 

342 else: 

343 object_key, field_path = None, fragment.split('.') 

344 

345 target = repaired 

346 if object_key is not None: 

347 target = repaired.setdefault(object_key, {}) 

348 if not isinstance(target, dict): 

349 continue 

350 

351 _assign_path(target, field_path, row.value) 

352 orphan_keys.append(row.key) 

353 repaired_any = True 

354 

355 if not repaired_any: 

356 continue 

357 

358 if existing: 

359 existing.value = repaired 

360 existing.updated_at = now 

361 else: 

362 db.add(Config(key=config_key, value=repaired, updated_at=now)) 

363 repaired_keys.append(config_key) 

364 

365 if orphan_keys: 

366 await db.execute(delete(Config).where(Config.key.in_(orphan_keys))) 

367 

368 for key in ('ui.default_models', 'ui.default_pinned_models'): 

369 row = await db.get(Config, key) 

370 if not row or not isinstance(row.value, list): 

371 continue 

372 

373 row.value = ','.join(model_id for model_id in (str(item).strip() for item in row.value) if model_id) 

374 row.updated_at = now 

375 default_model_keys.append(key) 

376 

377 if repaired_keys or orphan_keys or default_model_keys: 377 ↛ exitline 377 didn't jump to the function exit

378 await db.commit() 

379 if repaired_keys or orphan_keys: 

380 log.info('Repaired flattened dict config rows for %s', ', '.join(repaired_keys)) 

381 if default_model_keys: 

382 log.info('Repaired default model config rows for %s', ', '.join(default_model_keys))