mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-23 18:02:25 -05:00
Checking whether a user may reach a file loaded and validated every workspace model that user can access, then scanned each model's knowledge list in Python for one file id. Folder listings run that check once per file, so opening a folder of twenty files rebuilt the whole accessible-model set twenty times, and the same check sits on every retrieval and download path. The lookup now runs the other way round: the database returns the models that attach the file, and only those are access-checked. The text match on the metadata column is a prefilter and the knowledge entries still decide, so a file id that merely appears in a description grants nothing; file ids are server-generated uuids, so the match can only be too wide, never too narrow. Measured with 500 accessible workspace models: a single check drops from 9 queries and ~20 ms to 6 and ~2.6 ms, and a twenty-file folder listing from 180 queries and ~680 ms to 120 and ~56 ms. A 72-case matrix over owner, public, direct-user and group grants, for both read and write, returns exactly what it returned before, and write still requires the model owner to own the file. The check also no longer writes to the database while answering a read-only question.
651 lines
25 KiB
Python
Executable File
651 lines
25 KiB
Python
Executable File
from __future__ import annotations
|
|
|
|
import logging
|
|
import time
|
|
from copy import deepcopy
|
|
from typing import Any, Optional
|
|
|
|
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
|
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
|
|
from open_webui.models.groups import Groups
|
|
from open_webui.models.users import User, UserModel, UserResponse, Users
|
|
from open_webui.utils.misc import json_text_variants
|
|
from open_webui.utils.validate import validate_profile_image_url
|
|
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
|
from sqlalchemy import BigInteger, Boolean, Column, String, Text, cast, delete, func, or_, select, update
|
|
from sqlalchemy.dialects.postgresql import JSONB
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
log = logging.getLogger(__name__)
|
|
|
|
# Track invalid profile_image_url values we've already warned about so we
|
|
# don't flood the logs on every DB read (the validator fires per-row).
|
|
_warned_profile_urls: set[str] = set()
|
|
|
|
|
|
def strip_extracted_content_from_model_knowledge(knowledge: Any) -> Any:
|
|
"""Drop duplicated extracted text from ModelMeta.knowledge."""
|
|
if not isinstance(knowledge, list):
|
|
return knowledge
|
|
|
|
sanitized = []
|
|
|
|
for item in knowledge:
|
|
if not isinstance(item, dict):
|
|
sanitized.append(item)
|
|
continue
|
|
|
|
next_item = item
|
|
data = item.get('data')
|
|
if isinstance(data, dict) and 'content' in data:
|
|
next_item = deepcopy(item)
|
|
next_item.get('data', {}).pop('content', None)
|
|
|
|
file = next_item.get('file')
|
|
file_data = file.get('data') if isinstance(file, dict) else None
|
|
if isinstance(file_data, dict) and 'content' in file_data:
|
|
if next_item is item:
|
|
next_item = deepcopy(item)
|
|
file = next_item.get('file')
|
|
file_data = file.get('data') if isinstance(file, dict) else None
|
|
file_data.pop('content', None)
|
|
|
|
sanitized.append(next_item)
|
|
|
|
return sanitized
|
|
|
|
|
|
# --- Models DB Schema ---
|
|
|
|
|
|
class ModelParams(BaseModel):
|
|
"""Parameters for model inference (temperature, top_p, etc.)."""
|
|
|
|
model_config = ConfigDict(extra='allow')
|
|
|
|
|
|
class ModelMeta(BaseModel):
|
|
"""Metadata for a workspace model entry (profile, description, tags, capabilities)."""
|
|
|
|
profile_image_url: str | None = None
|
|
description: str | None = Field(default=None, description='User-facing description of the model.')
|
|
capabilities: dict | None = None
|
|
knowledge: list[Any] | None = None
|
|
|
|
model_config = ConfigDict(extra='allow')
|
|
|
|
@field_validator('profile_image_url', mode='before')
|
|
@classmethod
|
|
def check_profile_image_url(cls, v: str | None) -> str | None:
|
|
if v is None:
|
|
return v
|
|
try:
|
|
return validate_profile_image_url(v)
|
|
except ValueError:
|
|
if v not in _warned_profile_urls:
|
|
_warned_profile_urls.add(v)
|
|
log.warning(
|
|
'Clearing invalid profile_image_url stored in DB (likely a legacy SVG data-URI): %.80s…',
|
|
v,
|
|
)
|
|
return None
|
|
|
|
@field_validator('knowledge', mode='before')
|
|
@classmethod
|
|
def strip_knowledge_content(cls, v):
|
|
return strip_extracted_content_from_model_knowledge(v)
|
|
|
|
@model_validator(mode='before')
|
|
@classmethod
|
|
def normalize_tags(cls, data):
|
|
if isinstance(data, dict) and 'tags' in data:
|
|
raw_tags = data['tags']
|
|
if isinstance(raw_tags, list):
|
|
normalized = []
|
|
for tag in raw_tags:
|
|
if isinstance(tag, str):
|
|
normalized.append({'name': tag})
|
|
elif isinstance(tag, dict) and 'name' in tag:
|
|
normalized.append(tag)
|
|
data['tags'] = normalized
|
|
return data
|
|
|
|
|
|
class Model(Base):
|
|
"""Workspace model entry — wraps an upstream LLM with custom params and metadata."""
|
|
|
|
__tablename__ = 'model'
|
|
|
|
id = Column(Text, primary_key=True, unique=True) # API model identifier; overrides built-in when matching
|
|
user_id = Column(Text) # owner
|
|
base_model_id = Column(Text, nullable=True) # actual upstream model for proxied requests
|
|
name = Column(Text) # human-readable display name
|
|
params = Column(JSONField) # see ModelParams
|
|
meta = Column(JSONField) # see ModelMeta
|
|
is_active = Column(Boolean, default=True) # soft-disable toggle
|
|
updated_at = Column(BigInteger) # epoch seconds
|
|
created_at = Column(BigInteger) # epoch seconds
|
|
|
|
|
|
class ModelModel(BaseModel):
|
|
id: str
|
|
user_id: str
|
|
base_model_id: str | None = None
|
|
|
|
name: str
|
|
params: ModelParams
|
|
meta: ModelMeta
|
|
|
|
access_grants: list[AccessGrantModel] = Field(default_factory=list)
|
|
|
|
is_active: bool
|
|
updated_at: int # timestamp in epoch
|
|
created_at: int # timestamp in epoch
|
|
|
|
model_config = ConfigDict(
|
|
from_attributes=True,
|
|
)
|
|
|
|
|
|
class ModelUserResponse(ModelModel):
|
|
user: UserResponse | None = None
|
|
|
|
|
|
class ModelAccessResponse(ModelUserResponse):
|
|
write_access: bool | None = False
|
|
|
|
|
|
class ModelResponse(ModelModel):
|
|
pass
|
|
|
|
|
|
class ModelListResponse(BaseModel):
|
|
items: list[ModelUserResponse]
|
|
total: int
|
|
|
|
|
|
class ModelAccessListResponse(BaseModel):
|
|
items: list[ModelAccessResponse]
|
|
total: int
|
|
|
|
|
|
class ModelForm(BaseModel):
|
|
model_config = ConfigDict(extra='ignore')
|
|
|
|
id: str
|
|
base_model_id: str | None = None
|
|
name: str
|
|
meta: ModelMeta
|
|
params: ModelParams
|
|
access_grants: list[dict | None] = None
|
|
is_active: bool = True
|
|
|
|
|
|
class ModelsTable:
|
|
async def _get_access_grants(self, model_id: str, db: AsyncSession | None = None) -> list[AccessGrantModel]:
|
|
return await AccessGrants.get_grants_by_resource('model', model_id, db=db)
|
|
|
|
async def _to_model_model(
|
|
self,
|
|
model: Model,
|
|
access_grants: list[AccessGrantModel | None] = None,
|
|
db: AsyncSession | None = None,
|
|
) -> ModelModel:
|
|
if isinstance(model.meta, dict):
|
|
knowledge = model.meta.get('knowledge')
|
|
stripped_knowledge = strip_extracted_content_from_model_knowledge(knowledge)
|
|
if stripped_knowledge != knowledge:
|
|
model.meta = {**model.meta, 'knowledge': stripped_knowledge}
|
|
if db is not None:
|
|
await db.commit()
|
|
|
|
model_model = ModelModel.model_validate(model)
|
|
model_model.access_grants = (
|
|
access_grants if access_grants is not None else await self._get_access_grants(model_model.id, db=db)
|
|
)
|
|
return model_model
|
|
|
|
async def insert_new_model(
|
|
self, form_data: ModelForm, user_id: str, db: AsyncSession | None = None
|
|
) -> ModelModel | None:
|
|
try:
|
|
async with get_async_db_context(db) as db:
|
|
result = Model(
|
|
**{
|
|
**form_data.model_dump(exclude={'access_grants'}),
|
|
'user_id': user_id,
|
|
'created_at': int(time.time()),
|
|
'updated_at': int(time.time()),
|
|
}
|
|
)
|
|
db.add(result)
|
|
await db.commit()
|
|
await AccessGrants.set_access_grants('model', result.id, form_data.access_grants, db=db)
|
|
|
|
if result:
|
|
return await self._to_model_model(result, db=db)
|
|
else:
|
|
return None
|
|
except Exception as e:
|
|
log.exception(f'Failed to insert a new model: {e}')
|
|
return None
|
|
|
|
async def get_all_models(self, db: AsyncSession | None = None) -> list[ModelModel]:
|
|
async with get_async_db_context(db) as db:
|
|
result = await db.execute(select(Model))
|
|
all_models = result.scalars().all()
|
|
model_ids = [model.id for model in all_models]
|
|
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
|
models: list[ModelModel] = []
|
|
for model in all_models:
|
|
try:
|
|
models.append(await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db))
|
|
except Exception as exc:
|
|
log.error('Skipping model %r during get_all_models due to error: %s', model.id, exc)
|
|
return models
|
|
|
|
async def get_models(self, db: AsyncSession | None = None) -> list[ModelUserResponse]:
|
|
async with get_async_db_context(db) as db:
|
|
result = await db.execute(select(Model).filter(Model.base_model_id != None))
|
|
all_models = result.scalars().all()
|
|
|
|
user_ids = list(set(model.user_id for model in all_models))
|
|
model_ids = [model.id for model in all_models]
|
|
|
|
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
|
|
users_dict = {user.id: user for user in users}
|
|
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
|
|
|
models = []
|
|
for model in all_models:
|
|
user = users_dict.get(model.user_id)
|
|
models.append(
|
|
ModelUserResponse.model_validate(
|
|
{
|
|
**(
|
|
await self._to_model_model(
|
|
model,
|
|
access_grants=grants_map.get(model.id, []),
|
|
db=db,
|
|
)
|
|
).model_dump(),
|
|
'user': user.model_dump() if user else None,
|
|
}
|
|
)
|
|
)
|
|
return models
|
|
|
|
async def get_model_owners_attaching_file(self, file_id: str, db: AsyncSession | None = None) -> dict[str, str]:
|
|
"""Map of model id to owner id for workspace models whose knowledge attaches this file."""
|
|
async with get_async_db_context(db) as db:
|
|
# File ids are server-generated uuids, so the text match can only over-match.
|
|
result = await db.execute(
|
|
select(Model.id, Model.user_id, Model.meta).filter(
|
|
Model.base_model_id.is_not(None), cast(Model.meta, String).like(f'%"{file_id}"%')
|
|
)
|
|
)
|
|
return {
|
|
model_id: user_id
|
|
for model_id, user_id, meta in result.all()
|
|
if any(
|
|
isinstance(item, dict) and item.get('type') == 'file' and item.get('id') == file_id
|
|
for item in meta.get('knowledge') or []
|
|
)
|
|
}
|
|
|
|
@staticmethod
|
|
def _meta_has_tag(meta: dict | None, tag: str) -> bool:
|
|
if not meta:
|
|
return False
|
|
|
|
for raw_tag in meta.get('tags', []):
|
|
name = raw_tag.get('name') if isinstance(raw_tag, dict) else str(raw_tag)
|
|
if name == tag:
|
|
return True
|
|
|
|
return False
|
|
|
|
async def get_base_models(self, tag: str | None = None, db: AsyncSession | None = None) -> list[ModelModel]:
|
|
async with get_async_db_context(db) as db:
|
|
result = await db.execute(select(Model).filter(Model.base_model_id.is_(None)))
|
|
all_models = result.scalars().all()
|
|
if tag:
|
|
all_models = [model for model in all_models if self._meta_has_tag(model.meta, tag)]
|
|
|
|
model_ids = [model.id for model in all_models]
|
|
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
|
return [
|
|
await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db)
|
|
for model in all_models
|
|
]
|
|
|
|
async def get_models_by_user_id(self, user_id: str, db: AsyncSession | None = None) -> list[ModelUserResponse]:
|
|
models = await self.get_models(db=db)
|
|
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user_id, db=db)}
|
|
|
|
# One grants query for all non-owned models instead of one per model
|
|
accessible_ids = await AccessGrants.get_accessible_resource_ids(
|
|
user_id=user_id,
|
|
resource_type='model',
|
|
resource_ids=[model.id for model in models if model.user_id != user_id],
|
|
permission='write',
|
|
user_group_ids=user_group_ids,
|
|
db=db,
|
|
)
|
|
return [model for model in models if model.user_id == user_id or model.id in accessible_ids]
|
|
|
|
def _has_permission(self, db, query, filter: dict, permission: str = 'read'):
|
|
return AccessGrants.has_permission_filter(
|
|
db=db,
|
|
query=query,
|
|
DocumentModel=Model,
|
|
filter=filter,
|
|
resource_type='model',
|
|
permission=permission,
|
|
)
|
|
|
|
async def search_models(
|
|
self,
|
|
user_id: str,
|
|
filter: dict = {},
|
|
skip: int = 0,
|
|
limit: int = 30,
|
|
db: AsyncSession | None = None,
|
|
) -> ModelListResponse:
|
|
async with get_async_db_context(db) as db:
|
|
stmt = select(Model, User).outerjoin(User, User.id == Model.user_id)
|
|
stmt = stmt.filter(Model.base_model_id != None)
|
|
|
|
if filter:
|
|
query_key = filter.get('query')
|
|
if query_key:
|
|
stmt = stmt.filter(
|
|
or_(
|
|
Model.name.ilike(f'%{query_key}%'),
|
|
Model.base_model_id.ilike(f'%{query_key}%'),
|
|
User.name.ilike(f'%{query_key}%'),
|
|
User.email.ilike(f'%{query_key}%'),
|
|
User.username.ilike(f'%{query_key}%'),
|
|
)
|
|
)
|
|
|
|
view_option = filter.get('view_option')
|
|
if view_option == 'created':
|
|
stmt = stmt.filter(Model.user_id == user_id)
|
|
elif view_option == 'shared':
|
|
stmt = stmt.filter(Model.user_id != user_id)
|
|
|
|
# Apply access control filtering
|
|
stmt = self._has_permission(
|
|
db,
|
|
stmt,
|
|
filter,
|
|
permission='read',
|
|
)
|
|
|
|
tag = filter.get('tag')
|
|
if tag:
|
|
if db.bind.dialect.name == 'sqlite' and not tag.isascii():
|
|
# SQLite's LOWER() is ASCII-only, so match non-ASCII tags exact-case.
|
|
meta_text = cast(Model.meta, String)
|
|
variants = json_text_variants(tag)
|
|
else:
|
|
meta_text = func.lower(cast(Model.meta, String))
|
|
variants = json_text_variants(tag.lower())
|
|
stmt = stmt.filter(or_(*(meta_text.like(f'%"{variant}"%') for variant in variants)))
|
|
|
|
order_by = filter.get('order_by')
|
|
direction = filter.get('direction')
|
|
|
|
if order_by == 'name':
|
|
if direction == 'asc':
|
|
stmt = stmt.order_by(Model.name.asc())
|
|
else:
|
|
stmt = stmt.order_by(Model.name.desc())
|
|
elif order_by == 'created_at':
|
|
if direction == 'asc':
|
|
stmt = stmt.order_by(Model.created_at.asc())
|
|
else:
|
|
stmt = stmt.order_by(Model.created_at.desc())
|
|
elif order_by == 'updated_at':
|
|
if direction == 'asc':
|
|
stmt = stmt.order_by(Model.updated_at.asc())
|
|
else:
|
|
stmt = stmt.order_by(Model.updated_at.desc())
|
|
|
|
else:
|
|
stmt = stmt.order_by(Model.created_at.desc())
|
|
|
|
# Count BEFORE pagination
|
|
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
|
total = count_result.scalar()
|
|
|
|
if skip:
|
|
stmt = stmt.offset(skip)
|
|
if limit:
|
|
stmt = stmt.limit(limit)
|
|
|
|
result = await db.execute(stmt)
|
|
items = result.all()
|
|
|
|
model_ids = [model.id for model, _ in items]
|
|
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
|
|
|
models = []
|
|
for model, user in items:
|
|
models.append(
|
|
ModelUserResponse(
|
|
**(
|
|
await self._to_model_model(
|
|
model,
|
|
access_grants=grants_map.get(model.id, []),
|
|
db=db,
|
|
)
|
|
).model_dump(),
|
|
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
|
|
)
|
|
)
|
|
|
|
return ModelListResponse(items=models, total=total)
|
|
|
|
async def get_model_meta_by_id(self, id: str, db: AsyncSession | None = None) -> tuple[dict, int | None]:
|
|
"""Return (meta, updated_at) for a model, skipping access grant resolution."""
|
|
try:
|
|
async with get_async_db_context(db) as db:
|
|
result = await db.execute(select(Model.meta, Model.updated_at).filter_by(id=id))
|
|
return result.first()
|
|
except Exception:
|
|
return None
|
|
|
|
async def get_all_tags(
|
|
self,
|
|
user_id: str,
|
|
is_admin: bool = False,
|
|
is_base_model: bool = False,
|
|
db: AsyncSession | None = None,
|
|
) -> set[str]:
|
|
"""Extract unique tag names from model meta, querying only the meta column."""
|
|
async with get_async_db_context(db) as db:
|
|
stmt = select(Model.meta).filter(
|
|
Model.base_model_id.is_(None) if is_base_model else Model.base_model_id.is_not(None)
|
|
)
|
|
|
|
if not is_admin:
|
|
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
|
user_group_ids = [group.id for group in user_groups]
|
|
|
|
filter_dict = {'user_id': user_id}
|
|
if user_group_ids:
|
|
filter_dict['group_ids'] = user_group_ids
|
|
|
|
stmt = self._has_permission(db, stmt, filter_dict, permission='read')
|
|
|
|
result = await db.execute(stmt)
|
|
rows = result.scalars().all()
|
|
|
|
tags_set: set[str] = set()
|
|
for meta in rows:
|
|
if not meta:
|
|
continue
|
|
for tag in meta.get('tags', []):
|
|
try:
|
|
name = tag.get('name') if isinstance(tag, dict) else str(tag)
|
|
if name:
|
|
tags_set.add(name)
|
|
except Exception:
|
|
continue
|
|
|
|
return tags_set
|
|
|
|
async def get_model_by_id(self, id: str, db: AsyncSession | None = None) -> ModelModel | None:
|
|
try:
|
|
async with get_async_db_context(db) as db:
|
|
model = await db.get(Model, id)
|
|
return await self._to_model_model(model, db=db) if model else None
|
|
except Exception:
|
|
return None
|
|
|
|
async def get_models_by_ids(self, ids: list[str], db: AsyncSession | None = None) -> list[ModelModel]:
|
|
try:
|
|
async with get_async_db_context(db) as db:
|
|
result = await db.execute(select(Model).filter(Model.id.in_(ids)))
|
|
models = result.scalars().all()
|
|
model_ids = [model.id for model in models]
|
|
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
|
return [
|
|
await self._to_model_model(
|
|
model,
|
|
access_grants=grants_map.get(model.id, []),
|
|
db=db,
|
|
)
|
|
for model in models
|
|
]
|
|
except Exception:
|
|
return []
|
|
|
|
async def toggle_model_by_id(self, id: str, db: AsyncSession | None = None) -> ModelModel | None:
|
|
async with get_async_db_context(db) as db:
|
|
try:
|
|
result = await db.execute(select(Model).filter_by(id=id))
|
|
model = result.scalars().first()
|
|
if not model:
|
|
return None
|
|
|
|
model.is_active = not model.is_active
|
|
model.updated_at = int(time.time())
|
|
await db.commit()
|
|
|
|
return await self._to_model_model(model, db=db)
|
|
except Exception:
|
|
return None
|
|
|
|
async def update_model_by_id(self, id: str, model: ModelForm, db: AsyncSession | None = None) -> ModelModel | None:
|
|
try:
|
|
async with get_async_db_context(db) as db:
|
|
# update only the fields that are present in the model
|
|
data = model.model_dump(exclude={'id', 'access_grants'})
|
|
data['updated_at'] = int(time.time())
|
|
await db.execute(update(Model).filter_by(id=id).values(**data))
|
|
|
|
await db.commit()
|
|
if model.access_grants is not None:
|
|
await AccessGrants.set_access_grants('model', id, model.access_grants, db=db)
|
|
|
|
return await self.get_model_by_id(id, db=db)
|
|
except Exception as e:
|
|
log.exception(f'Failed to update the model by id {id}: {e}')
|
|
return None
|
|
|
|
async def update_model_updated_at_by_id(self, id: str, db: AsyncSession | None = None) -> ModelModel | None:
|
|
try:
|
|
async with get_async_db_context(db) as db:
|
|
result = await db.execute(select(Model).filter_by(id=id))
|
|
model = result.scalars().first()
|
|
if not model:
|
|
return None
|
|
model.updated_at = int(time.time())
|
|
await db.commit()
|
|
return await self._to_model_model(model, db=db)
|
|
except Exception as e:
|
|
log.exception(f'Failed to update the model updated_at by id {id}: {e}')
|
|
return None
|
|
|
|
async def delete_model_by_id(self, id: str, db: AsyncSession | None = None) -> bool:
|
|
try:
|
|
async with get_async_db_context(db) as db:
|
|
await AccessGrants.revoke_all_access('model', id, db=db)
|
|
await db.execute(delete(Model).filter_by(id=id))
|
|
await db.commit()
|
|
|
|
return True
|
|
except Exception:
|
|
return False
|
|
|
|
async def delete_all_models(self, db: AsyncSession | None = None) -> bool:
|
|
try:
|
|
async with get_async_db_context(db) as db:
|
|
result = await db.execute(select(Model.id))
|
|
model_ids = [row[0] for row in result.all()]
|
|
for model_id in model_ids:
|
|
await AccessGrants.revoke_all_access('model', model_id, db=db)
|
|
await db.execute(delete(Model))
|
|
await db.commit()
|
|
|
|
return True
|
|
except Exception:
|
|
return False
|
|
|
|
async def sync_models(
|
|
self, user_id: str, models: list[ModelModel], db: AsyncSession | None = None
|
|
) -> list[ModelModel]:
|
|
try:
|
|
async with get_async_db_context(db) as db:
|
|
# Get existing models
|
|
result = await db.execute(select(Model))
|
|
existing_models = result.scalars().all()
|
|
existing_ids = {model.id for model in existing_models}
|
|
|
|
# Prepare a set of new model IDs
|
|
new_model_ids = {model.id for model in models}
|
|
|
|
# Update or insert models
|
|
for model in models:
|
|
model_data = {
|
|
**model.model_dump(exclude={'access_grants'}),
|
|
'user_id': user_id,
|
|
'updated_at': int(time.time()),
|
|
}
|
|
|
|
if model.id in existing_ids:
|
|
await db.execute(update(Model).filter_by(id=model.id).values(**model_data))
|
|
else:
|
|
db.add(Model(**model_data))
|
|
await AccessGrants.set_access_grants('model', model.id, model.access_grants, db=db)
|
|
|
|
# Remove models that are no longer present
|
|
for model in existing_models:
|
|
if model.id not in new_model_ids:
|
|
await AccessGrants.revoke_all_access('model', model.id, db=db)
|
|
await db.delete(model)
|
|
|
|
await db.commit()
|
|
|
|
result = await db.execute(select(Model))
|
|
all_models = result.scalars().all()
|
|
model_ids = [model.id for model in all_models]
|
|
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
|
return [
|
|
await self._to_model_model(
|
|
model,
|
|
access_grants=grants_map.get(model.id, []),
|
|
db=db,
|
|
)
|
|
for model in all_models
|
|
]
|
|
except Exception as e:
|
|
log.exception(f'Error syncing models for user {user_id}: {e}')
|
|
return []
|
|
|
|
|
|
Models = ModelsTable() # singleton model registry
|