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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1from __future__ import annotations
3import logging
4import time
5from copy import deepcopy
6from typing import Any
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
18log = logging.getLogger(__name__)
21def normalize_model_tags(tags: Any) -> list[dict[str, str]]:
22 if not isinstance(tags, list):
23 return []
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
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
38 sanitized = []
40 for item in knowledge:
41 if not isinstance(item, dict):
42 sanitized.append(item)
43 continue
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)
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)
60 sanitized.append(next_item)
62 return sanitized
65# --- Models DB Schema ---
68class ModelParams(BaseModel):
69 """Parameters for model inference (temperature, top_p, etc.)."""
71 model_config = ConfigDict(extra='allow')
74class ModelMeta(BaseModel):
75 """Metadata for a workspace model entry (profile, description, tags, capabilities)."""
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
84 model_config = ConfigDict(extra='allow')
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
98 @field_validator('knowledge', mode='before')
99 @classmethod
100 def strip_knowledge_content(cls, v):
101 return strip_extracted_content_from_model_knowledge(v)
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
111class Model(Base):
112 """Workspace model entry — wraps an upstream LLM with custom params and metadata."""
114 __tablename__ = 'model'
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
127class ModelModel(BaseModel):
128 id: str
129 user_id: str
130 base_model_id: str | None = None
132 name: str
133 params: ModelParams
134 meta: ModelMeta
136 access_grants: list[AccessGrantModel] = Field(default_factory=list)
138 is_active: bool
139 updated_at: int # timestamp in epoch
140 created_at: int # timestamp in epoch
142 model_config = ConfigDict(
143 from_attributes=True,
144 )
147class ModelUserResponse(ModelModel):
148 user: UserResponse | None = None
151class ModelAccessResponse(ModelUserResponse):
152 write_access: bool | None = False
155class ModelResponse(ModelModel):
156 pass
159class ModelListResponse(BaseModel):
160 items: list[ModelUserResponse]
161 total: int
164class ModelAccessListResponse(BaseModel):
165 items: list[ModelAccessResponse]
166 total: int
169class ModelForm(BaseModel):
170 model_config = ConfigDict(extra='ignore')
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
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)
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()
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
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)
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
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
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)
250 if ids is not None:
251 stmt = stmt.filter(Model.id.in_(ids))
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 )
261 result = await db.execute(stmt)
262 all_models = result.scalars().all()
264 user_ids = list(set(model.user_id for model in all_models))
265 model_ids = [model.id for model in all_models]
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)
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
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 }
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
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
321 return False
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)]
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 ]
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 )
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)
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 )
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)
378 # Apply access control filtering
379 stmt = self._has_permission(
380 db,
381 stmt,
382 filter,
383 permission='read',
384 )
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)))
397 order_by = filter.get('order_by')
398 direction = filter.get('direction')
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())
416 else:
417 stmt = stmt.order_by(Model.created_at.desc())
419 # Count BEFORE pagination
420 count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
421 total = count_result.scalar()
423 if skip:
424 stmt = stmt.offset(skip)
425 if limit:
426 stmt = stmt.limit(limit)
428 result = await db.execute(stmt)
429 items = result.all()
431 model_ids = [model.id for model, _ in items]
432 grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
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 )
449 return ModelListResponse(items=models, total=total)
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
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 )
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]
479 filter_dict = {'user_id': user_id}
480 if user_group_ids:
481 filter_dict['group_ids'] = user_group_ids
483 stmt = self._has_permission(db, stmt, filter_dict, permission='read')
485 result = await db.execute(stmt)
486 rows = result.scalars().all()
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
500 return tags_set
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
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 []
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
536 model.is_active = not model.is_active
537 model.updated_at = int(time.time())
538 await db.commit()
540 return await self._to_model_model(model, db=db)
541 except Exception:
542 return None
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))
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)
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
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
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()
582 return True
583 except Exception:
584 return False
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()
596 return True
597 except Exception:
598 return False
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}
610 # Prepare a set of new model IDs
611 new_model_ids = {model.id for model in models}
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 }
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)
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)
633 await db.commit()
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 []
652Models = ModelsTable() # singleton model registry