Coverage for open_webui/models/models.py: 66%

392 statements  

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

1from __future__ import annotations 

2 

3import logging 

4import time 

5from copy import deepcopy 

6from typing import Any 

7 

8from open_webui.internal.db import Base, JSONField, get_async_db_context 

9from open_webui.models.access_grants import AccessGrantModel, AccessGrants 

10from open_webui.models.groups import Groups 

11from open_webui.models.users import User, UserModel, UserResponse, Users 

12from open_webui.utils.misc import json_text_variants 

13from open_webui.utils.validate import validate_image_url 

14from pydantic import BaseModel, ConfigDict, Field, ValidationInfo, field_validator, model_validator 

15from sqlalchemy import BigInteger, Boolean, Column, String, Text, cast, delete, func, or_, select, update 

16from sqlalchemy.ext.asyncio import AsyncSession 

17 

18log = logging.getLogger(__name__) 

19 

20 

21def normalize_model_tags(tags: Any) -> list[dict[str, str]]: 

22 if not isinstance(tags, list): 

23 return [] 

24 

25 normalized = [] 

26 for tag in tags: 

27 name = tag.get('name') if isinstance(tag, dict) else tag 

28 if isinstance(name, str) and name.strip(): 

29 normalized.append({'name': name.strip()}) 

30 return normalized 

31 

32 

33def strip_extracted_content_from_model_knowledge(knowledge: Any) -> Any: 

34 """Drop duplicated extracted text from ModelMeta.knowledge.""" 

35 if not isinstance(knowledge, list): 

36 return knowledge 

37 

38 sanitized = [] 

39 

40 for item in knowledge: 

41 if not isinstance(item, dict): 

42 sanitized.append(item) 

43 continue 

44 

45 next_item = item 

46 data = item.get('data') 

47 if isinstance(data, dict) and 'content' in data: 47 ↛ 48line 47 didn't jump to line 48 because the condition on line 47 was never true

48 next_item = deepcopy(item) 

49 next_item.get('data', {}).pop('content', None) 

50 

51 file = next_item.get('file') 

52 file_data = file.get('data') if isinstance(file, dict) else None 

53 if isinstance(file_data, dict) and 'content' in file_data: 53 ↛ 54line 53 didn't jump to line 54 because the condition on line 53 was never true

54 if next_item is item: 

55 next_item = deepcopy(item) 

56 file = next_item.get('file') 

57 file_data = file.get('data') if isinstance(file, dict) else None 

58 file_data.pop('content', None) 

59 

60 sanitized.append(next_item) 

61 

62 return sanitized 

63 

64 

65# --- Models DB Schema --- 

66 

67 

68class ModelParams(BaseModel): 

69 """Parameters for model inference (temperature, top_p, etc.).""" 

70 

71 model_config = ConfigDict(extra='allow') 

72 

73 

74class ModelMeta(BaseModel): 

75 """Metadata for a workspace model entry (profile, description, tags, capabilities).""" 

76 

77 profile_image_url: str | None = None 

78 background_image_url: str | None = None 

79 description: str | None = Field(default=None, description='User-facing description of the model.') 

80 i18n: dict[str, Any] | None = None 

81 capabilities: dict | None = None 

82 knowledge: list[Any] | None = None 

83 

84 model_config = ConfigDict(extra='allow') 

85 

86 @field_validator('profile_image_url', 'background_image_url', mode='before') 

87 @classmethod 

88 def check_image_url(cls, v: str | None, info: ValidationInfo) -> str | None: 

89 if v is None: 

90 return v 

91 try: 

92 return validate_image_url(v, file_only=info.field_name == 'background_image_url') 

93 except ValueError: 

94 if info.field_name == 'background_image_url': 

95 raise 

96 return None 

97 

98 @field_validator('knowledge', mode='before') 

99 @classmethod 

100 def strip_knowledge_content(cls, v): 

101 return strip_extracted_content_from_model_knowledge(v) 

102 

103 @model_validator(mode='before') 

104 @classmethod 

105 def normalize_tags(cls, data): 

106 if isinstance(data, dict) and 'tags' in data: 106 ↛ 107line 106 didn't jump to line 107 because the condition on line 106 was never true

107 data['tags'] = normalize_model_tags(data['tags']) 

108 return data 

109 

110 

111class Model(Base): 

112 """Workspace model entry — wraps an upstream LLM with custom params and metadata.""" 

113 

114 __tablename__ = 'model' 

115 

116 id = Column(Text, primary_key=True, unique=True) # API model identifier; overrides built-in when matching 

117 user_id = Column(Text) # owner 

118 base_model_id = Column(Text, nullable=True) # actual upstream model for proxied requests 

119 name = Column(Text) # human-readable display name 

120 params = Column(JSONField) # see ModelParams 

121 meta = Column(JSONField) # see ModelMeta 

122 is_active = Column(Boolean, default=True) # soft-disable toggle 

123 updated_at = Column(BigInteger) # epoch seconds 

124 created_at = Column(BigInteger) # epoch seconds 

125 

126 

127class ModelModel(BaseModel): 

128 id: str 

129 user_id: str 

130 base_model_id: str | None = None 

131 

132 name: str 

133 params: ModelParams 

134 meta: ModelMeta 

135 

136 access_grants: list[AccessGrantModel] = Field(default_factory=list) 

137 

138 is_active: bool 

139 updated_at: int # timestamp in epoch 

140 created_at: int # timestamp in epoch 

141 

142 model_config = ConfigDict( 

143 from_attributes=True, 

144 ) 

145 

146 

147class ModelUserResponse(ModelModel): 

148 user: UserResponse | None = None 

149 

150 

151class ModelAccessResponse(ModelUserResponse): 

152 write_access: bool | None = False 

153 

154 

155class ModelResponse(ModelModel): 

156 pass 

157 

158 

159class ModelListResponse(BaseModel): 

160 items: list[ModelUserResponse] 

161 total: int 

162 

163 

164class ModelAccessListResponse(BaseModel): 

165 items: list[ModelAccessResponse] 

166 total: int 

167 

168 

169class ModelForm(BaseModel): 

170 model_config = ConfigDict(extra='ignore') 

171 

172 id: str = Field(pattern=r'^\S+$') 

173 base_model_id: str | None = None 

174 name: str 

175 meta: ModelMeta 

176 params: ModelParams 

177 access_grants: list[dict] | None = None 

178 is_active: bool = True 

179 

180 

181class ModelsTable: 

182 async def _get_access_grants(self, model_id: str, db: AsyncSession | None = None) -> list[AccessGrantModel]: 

183 return await AccessGrants.get_grants_by_resource('model', model_id, db=db) 

184 

185 async def _to_model_model( 

186 self, 

187 model: Model, 

188 access_grants: list[AccessGrantModel] | None = None, 

189 db: AsyncSession | None = None, 

190 ) -> ModelModel: 

191 if isinstance(model.meta, dict): 191 ↛ 199line 191 didn't jump to line 199 because the condition on line 191 was always true

192 knowledge = model.meta.get('knowledge') 

193 stripped_knowledge = strip_extracted_content_from_model_knowledge(knowledge) 

194 if stripped_knowledge != knowledge: 194 ↛ 195line 194 didn't jump to line 195 because the condition on line 194 was never true

195 model.meta = {**model.meta, 'knowledge': stripped_knowledge} 

196 if db is not None: 

197 await db.commit() 

198 

199 model_model = ModelModel.model_validate(model) 

200 model_model.access_grants = ( 

201 access_grants if access_grants is not None else await self._get_access_grants(model_model.id, db=db) 

202 ) 

203 return model_model 

204 

205 async def insert_new_model( 

206 self, form_data: ModelForm, user_id: str, db: AsyncSession | None = None 

207 ) -> ModelModel | None: 

208 try: 

209 async with get_async_db_context(db) as db: 

210 result = Model( 

211 **{ 

212 **form_data.model_dump(exclude={'access_grants'}), 

213 'user_id': user_id, 

214 'created_at': int(time.time()), 

215 'updated_at': int(time.time()), 

216 } 

217 ) 

218 db.add(result) 

219 await db.commit() 

220 await AccessGrants.set_access_grants('model', result.id, form_data.access_grants, db=db) 

221 

222 if result: 222 ↛ 225line 222 didn't jump to line 225 because the condition on line 222 was always true

223 return await self._to_model_model(result, db=db) 

224 else: 

225 return None 

226 except Exception as e: 

227 log.exception(f'Failed to insert a new model: {e}') 

228 return None 

229 

230 async def get_all_models(self, db: AsyncSession | None = None) -> list[ModelModel]: 

231 async with get_async_db_context(db) as db: 

232 result = await db.execute(select(Model)) 

233 all_models = result.scalars().all() 

234 model_ids = [model.id for model in all_models] 

235 grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) 

236 models: list[ModelModel] = [] 

237 for model in all_models: 

238 try: 

239 models.append(await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db)) 

240 except Exception as exc: 

241 log.error('Skipping model %r during get_all_models due to error: %s', model.id, exc) 

242 return models 

243 

244 async def get_models( 

245 self, writable_by_user_id: str | None = None, db: AsyncSession | None = None, ids: list[str] | None = None 

246 ) -> list[ModelUserResponse]: 

247 async with get_async_db_context(db) as db: 

248 stmt = select(Model).filter(Model.base_model_id != None) 

249 

250 if ids is not None: 

251 stmt = stmt.filter(Model.id.in_(ids)) 

252 

253 if writable_by_user_id: 253 ↛ 254line 253 didn't jump to line 254 because the condition on line 253 was never true

254 user_group_ids = { 

255 group.id for group in await Groups.get_groups_by_member_id(writable_by_user_id, db=db) 

256 } 

257 stmt = self._has_permission( 

258 db, stmt, {'user_id': writable_by_user_id, 'group_ids': user_group_ids}, permission='write' 

259 ) 

260 

261 result = await db.execute(stmt) 

262 all_models = result.scalars().all() 

263 

264 user_ids = list(set(model.user_id for model in all_models)) 

265 model_ids = [model.id for model in all_models] 

266 

267 users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] 

268 users_dict = {user.id: user for user in users} 

269 grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) 

270 

271 models = [] 

272 for model in all_models: 

273 user = users_dict.get(model.user_id) 

274 models.append( 

275 ModelUserResponse.model_validate( 

276 { 

277 **( 

278 await self._to_model_model( 

279 model, 

280 access_grants=grants_map.get(model.id, []), 

281 db=db, 

282 ) 

283 ).model_dump(), 

284 'user': user.model_dump() if user else None, 

285 } 

286 ) 

287 ) 

288 return models 

289 

290 async def get_model_owner_ids_by_file_id( 

291 self, file_id: str, db: AsyncSession | None = None, include_background: bool = False 

292 ) -> dict[str, str]: 

293 """Return model IDs mapped to owner IDs for models referencing the file.""" 

294 async with get_async_db_context(db) as db: 

295 # File ids are server-generated uuids, so the text match can only over-match. 

296 result = await db.execute( 

297 select(Model.id, Model.user_id, Model.meta).filter( 

298 Model.base_model_id.is_not(None), cast(Model.meta, String).like(f'%{file_id}%') 

299 ) 

300 ) 

301 return { 

302 model_id: user_id 

303 for model_id, user_id, meta in result.all() 

304 if any( 

305 isinstance(item, dict) and item.get('type') == 'file' and item.get('id') == file_id 

306 for item in meta.get('knowledge') or [] 

307 ) 

308 or (include_background and meta.get('background_image_url') == f'/api/v1/files/{file_id}/content') 

309 } 

310 

311 @staticmethod 

312 def _meta_has_tag(meta: dict | None, tag: str) -> bool: 

313 if not meta: 313 ↛ 314line 313 didn't jump to line 314 because the condition on line 313 was never true

314 return False 

315 

316 for raw_tag in meta.get('tags', []): 316 ↛ 317line 316 didn't jump to line 317 because the loop on line 316 never started

317 name = raw_tag.get('name') if isinstance(raw_tag, dict) else str(raw_tag) 

318 if name == tag: 

319 return True 

320 

321 return False 

322 

323 async def get_base_models(self, tag: str | None = None, db: AsyncSession | None = None) -> list[ModelModel]: 

324 async with get_async_db_context(db) as db: 

325 result = await db.execute(select(Model).filter(Model.base_model_id.is_(None))) 

326 all_models = result.scalars().all() 

327 if tag: 

328 all_models = [model for model in all_models if self._meta_has_tag(model.meta, tag)] 

329 

330 model_ids = [model.id for model in all_models] 

331 grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) 

332 return [ 

333 await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db) 

334 for model in all_models 

335 ] 

336 

337 def _has_permission(self, db, query, filter: dict, permission: str = 'read'): 

338 return AccessGrants.has_permission_filter( 

339 db=db, 

340 query=query, 

341 DocumentModel=Model, 

342 filter=filter, 

343 resource_type='model', 

344 permission=permission, 

345 ) 

346 

347 async def search_models( 

348 self, 

349 user_id: str, 

350 filter: dict = {}, 

351 skip: int = 0, 

352 limit: int = 30, 

353 db: AsyncSession | None = None, 

354 ) -> ModelListResponse: 

355 async with get_async_db_context(db) as db: 

356 stmt = select(Model, User).outerjoin(User, User.id == Model.user_id) 

357 stmt = stmt.filter(Model.base_model_id != None) 

358 

359 if filter: 

360 query_key = filter.get('query') 

361 if query_key: 

362 stmt = stmt.filter( 

363 or_( 

364 Model.name.ilike(f'%{query_key}%'), 

365 Model.base_model_id.ilike(f'%{query_key}%'), 

366 User.name.ilike(f'%{query_key}%'), 

367 User.email.ilike(f'%{query_key}%'), 

368 User.username.ilike(f'%{query_key}%'), 

369 ) 

370 ) 

371 

372 view_option = filter.get('view_option') 

373 if view_option == 'created': 373 ↛ 374line 373 didn't jump to line 374 because the condition on line 373 was never true

374 stmt = stmt.filter(Model.user_id == user_id) 

375 elif view_option == 'shared': 375 ↛ 376line 375 didn't jump to line 376 because the condition on line 375 was never true

376 stmt = stmt.filter(Model.user_id != user_id) 

377 

378 # Apply access control filtering 

379 stmt = self._has_permission( 

380 db, 

381 stmt, 

382 filter, 

383 permission='read', 

384 ) 

385 

386 tag = filter.get('tag') 

387 if tag: 

388 if db.bind.dialect.name == 'sqlite' and not tag.isascii(): 

389 # SQLite's LOWER() is ASCII-only, so match non-ASCII tags exact-case. 

390 meta_text = cast(Model.meta, String) 

391 variants = json_text_variants(tag) 

392 else: 

393 meta_text = func.lower(cast(Model.meta, String)) 

394 variants = json_text_variants(tag.lower()) 

395 stmt = stmt.filter(or_(*(meta_text.like(f'%"{variant}"%') for variant in variants))) 

396 

397 order_by = filter.get('order_by') 

398 direction = filter.get('direction') 

399 

400 if order_by == 'name': 400 ↛ 401line 400 didn't jump to line 401 because the condition on line 400 was never true

401 if direction == 'asc': 

402 stmt = stmt.order_by(Model.name.asc()) 

403 else: 

404 stmt = stmt.order_by(Model.name.desc()) 

405 elif order_by == 'created_at': 405 ↛ 406line 405 didn't jump to line 406 because the condition on line 405 was never true

406 if direction == 'asc': 

407 stmt = stmt.order_by(Model.created_at.asc()) 

408 else: 

409 stmt = stmt.order_by(Model.created_at.desc()) 

410 elif order_by == 'updated_at': 410 ↛ 411line 410 didn't jump to line 411 because the condition on line 410 was never true

411 if direction == 'asc': 

412 stmt = stmt.order_by(Model.updated_at.asc()) 

413 else: 

414 stmt = stmt.order_by(Model.updated_at.desc()) 

415 

416 else: 

417 stmt = stmt.order_by(Model.created_at.desc()) 

418 

419 # Count BEFORE pagination 

420 count_result = await db.execute(select(func.count()).select_from(stmt.subquery())) 

421 total = count_result.scalar() 

422 

423 if skip: 

424 stmt = stmt.offset(skip) 

425 if limit: 

426 stmt = stmt.limit(limit) 

427 

428 result = await db.execute(stmt) 

429 items = result.all() 

430 

431 model_ids = [model.id for model, _ in items] 

432 grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) 

433 

434 models = [] 

435 for model, user in items: 

436 models.append( 

437 ModelUserResponse( 

438 **( 

439 await self._to_model_model( 

440 model, 

441 access_grants=grants_map.get(model.id, []), 

442 db=db, 

443 ) 

444 ).model_dump(), 

445 user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None), 

446 ) 

447 ) 

448 

449 return ModelListResponse(items=models, total=total) 

450 

451 async def get_model_meta_by_id( 

452 self, id: str, db: AsyncSession | None = None 

453 ) -> tuple[dict, str, int | None] | None: 

454 """Return (meta, user_id, updated_at) for a model, skipping access grant resolution.""" 

455 try: 

456 async with get_async_db_context(db) as db: 

457 result = await db.execute(select(Model.meta, Model.user_id, Model.updated_at).filter_by(id=id)) 

458 return result.first() 

459 except Exception: 

460 return None 

461 

462 async def get_all_tags( 

463 self, 

464 user_id: str, 

465 is_admin: bool = False, 

466 is_base_model: bool = False, 

467 db: AsyncSession | None = None, 

468 ) -> set[str]: 

469 """Extract unique tag names from model meta, querying only the meta column.""" 

470 async with get_async_db_context(db) as db: 

471 stmt = select(Model.meta).filter( 

472 Model.base_model_id.is_(None) if is_base_model else Model.base_model_id.is_not(None) 

473 ) 

474 

475 if not is_admin: 475 ↛ 476line 475 didn't jump to line 476 because the condition on line 475 was never true

476 user_groups = await Groups.get_groups_by_member_id(user_id, db=db) 

477 user_group_ids = [group.id for group in user_groups] 

478 

479 filter_dict = {'user_id': user_id} 

480 if user_group_ids: 

481 filter_dict['group_ids'] = user_group_ids 

482 

483 stmt = self._has_permission(db, stmt, filter_dict, permission='read') 

484 

485 result = await db.execute(stmt) 

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

487 

488 tags_set: set[str] = set() 

489 for meta in rows: 

490 if not meta: 

491 continue 

492 for tag in meta.get('tags', []): 

493 try: 

494 name = tag.get('name') if isinstance(tag, dict) else str(tag) 

495 if name: 

496 tags_set.add(name) 

497 except Exception: 

498 continue 

499 

500 return tags_set 

501 

502 async def get_model_by_id(self, id: str, db: AsyncSession | None = None) -> ModelModel | None: 

503 try: 

504 async with get_async_db_context(db) as db: 

505 model = await db.get(Model, id) 

506 return await self._to_model_model(model, db=db) if model else None 

507 except Exception: 

508 return None 

509 

510 async def get_models_by_ids(self, ids: list[str], db: AsyncSession | None = None) -> list[ModelModel]: 

511 try: 

512 async with get_async_db_context(db) as db: 

513 result = await db.execute(select(Model).filter(Model.id.in_(ids))) 

514 models = result.scalars().all() 

515 model_ids = [model.id for model in models] 

516 grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) 

517 return [ 

518 await self._to_model_model( 

519 model, 

520 access_grants=grants_map.get(model.id, []), 

521 db=db, 

522 ) 

523 for model in models 

524 ] 

525 except Exception: 

526 return [] 

527 

528 async def toggle_model_by_id(self, id: str, db: AsyncSession | None = None) -> ModelModel | None: 

529 async with get_async_db_context(db) as db: 

530 try: 

531 result = await db.execute(select(Model).filter_by(id=id)) 

532 model = result.scalars().first() 

533 if not model: 

534 return None 

535 

536 model.is_active = not model.is_active 

537 model.updated_at = int(time.time()) 

538 await db.commit() 

539 

540 return await self._to_model_model(model, db=db) 

541 except Exception: 

542 return None 

543 

544 async def update_model_by_id(self, id: str, model: ModelForm, db: AsyncSession | None = None) -> ModelModel | None: 

545 try: 

546 async with get_async_db_context(db) as db: 

547 # update only the fields that are present in the model 

548 data = model.model_dump(exclude={'id', 'access_grants'}) 

549 data['updated_at'] = int(time.time()) 

550 await db.execute(update(Model).filter_by(id=id).values(**data)) 

551 

552 await db.commit() 

553 if model.access_grants is not None: 

554 await AccessGrants.set_access_grants('model', id, model.access_grants, db=db) 

555 

556 return await self.get_model_by_id(id, db=db) 

557 except Exception as e: 

558 log.exception(f'Failed to update the model by id {id}: {e}') 

559 return None 

560 

561 async def update_model_updated_at_by_id(self, id: str, db: AsyncSession | None = None) -> ModelModel | None: 

562 try: 

563 async with get_async_db_context(db) as db: 

564 result = await db.execute(select(Model).filter_by(id=id)) 

565 model = result.scalars().first() 

566 if not model: 

567 return None 

568 model.updated_at = int(time.time()) 

569 await db.commit() 

570 return await self._to_model_model(model, db=db) 

571 except Exception as e: 

572 log.exception(f'Failed to update the model updated_at by id {id}: {e}') 

573 return None 

574 

575 async def delete_model_by_id(self, id: str, db: AsyncSession | None = None) -> bool: 

576 try: 

577 async with get_async_db_context(db) as db: 

578 await AccessGrants.revoke_all_access('model', id, db=db) 

579 await db.execute(delete(Model).filter_by(id=id)) 

580 await db.commit() 

581 

582 return True 

583 except Exception: 

584 return False 

585 

586 async def delete_all_models(self, db: AsyncSession | None = None) -> bool: 

587 try: 

588 async with get_async_db_context(db) as db: 

589 result = await db.execute(select(Model.id)) 

590 model_ids = [row[0] for row in result.all()] 

591 for model_id in model_ids: 

592 await AccessGrants.revoke_all_access('model', model_id, db=db) 

593 await db.execute(delete(Model)) 

594 await db.commit() 

595 

596 return True 

597 except Exception: 

598 return False 

599 

600 async def sync_models( 

601 self, user_id: str, models: list[ModelModel], db: AsyncSession | None = None 

602 ) -> list[ModelModel]: 

603 try: 

604 async with get_async_db_context(db) as db: 

605 # Get existing models 

606 result = await db.execute(select(Model)) 

607 existing_models = result.scalars().all() 

608 existing_ids = {model.id for model in existing_models} 

609 

610 # Prepare a set of new model IDs 

611 new_model_ids = {model.id for model in models} 

612 

613 # Update or insert models 

614 for model in models: 

615 model_data = { 

616 **model.model_dump(exclude={'access_grants'}), 

617 'user_id': user_id, 

618 'updated_at': int(time.time()), 

619 } 

620 

621 if model.id in existing_ids: 621 ↛ 622line 621 didn't jump to line 622 because the condition on line 621 was never true

622 await db.execute(update(Model).filter_by(id=model.id).values(**model_data)) 

623 else: 

624 db.add(Model(**model_data)) 

625 await AccessGrants.set_access_grants('model', model.id, model.access_grants, db=db) 

626 

627 # Remove models that are no longer present 

628 for model in existing_models: 

629 if model.id not in new_model_ids: 629 ↛ 628line 629 didn't jump to line 628 because the condition on line 629 was always true

630 await AccessGrants.revoke_all_access('model', model.id, db=db) 

631 await db.delete(model) 

632 

633 await db.commit() 

634 

635 result = await db.execute(select(Model)) 

636 all_models = result.scalars().all() 

637 model_ids = [model.id for model in all_models] 

638 grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db) 

639 return [ 

640 await self._to_model_model( 

641 model, 

642 access_grants=grants_map.get(model.id, []), 

643 db=db, 

644 ) 

645 for model in all_models 

646 ] 

647 except Exception as e: 

648 log.exception(f'Error syncing models for user {user_id}: {e}') 

649 return [] 

650 

651 

652Models = ModelsTable() # singleton model registry