Skip to content

Commit 8c485b2

Browse files
committed
refac
1 parent d664922 commit 8c485b2

8 files changed

Lines changed: 271 additions & 38 deletions

File tree

backend/open_webui/models/access_grants.py

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -456,6 +456,31 @@ def get_grants_by_resource(
456456
)
457457
return [AccessGrantModel.model_validate(g) for g in grants]
458458

459+
def get_grants_by_resources(
460+
self,
461+
resource_type: str,
462+
resource_ids: list[str],
463+
db: Optional[Session] = None,
464+
) -> dict[str, list[AccessGrantModel]]:
465+
"""Batch-fetch grants for multiple resources. Returns {resource_id: [grants]}."""
466+
if not resource_ids:
467+
return {}
468+
with get_db_context(db) as db:
469+
grants = (
470+
db.query(AccessGrant)
471+
.filter(
472+
AccessGrant.resource_type == resource_type,
473+
AccessGrant.resource_id.in_(resource_ids),
474+
)
475+
.all()
476+
)
477+
result: dict[str, list[AccessGrantModel]] = {
478+
rid: [] for rid in resource_ids
479+
}
480+
for g in grants:
481+
result[g.resource_id].append(AccessGrantModel.model_validate(g))
482+
return result
483+
459484
def has_access(
460485
self,
461486
user_id: str,

backend/open_webui/models/channels.py

Lines changed: 42 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -261,13 +261,19 @@ def _get_access_grants(
261261
return AccessGrants.get_grants_by_resource("channel", channel_id, db=db)
262262

263263
def _to_channel_model(
264-
self, channel: Channel, db: Optional[Session] = None
264+
self,
265+
channel: Channel,
266+
access_grants: Optional[list[AccessGrantModel]] = None,
267+
db: Optional[Session] = None,
265268
) -> ChannelModel:
266269
channel_data = ChannelModel.model_validate(channel).model_dump(
267270
exclude={"access_grants"}
268271
)
269-
access_grants = self._get_access_grants(channel_data["id"], db=db)
270-
channel_data["access_grants"] = access_grants
272+
channel_data["access_grants"] = (
273+
access_grants
274+
if access_grants is not None
275+
else self._get_access_grants(channel_data["id"], db=db)
276+
)
271277
return ChannelModel.model_validate(channel_data)
272278

273279
def _collect_unique_user_ids(
@@ -368,7 +374,18 @@ def insert_new_channel(
368374
def get_channels(self, db: Optional[Session] = None) -> list[ChannelModel]:
369375
with get_db_context(db) as db:
370376
channels = db.query(Channel).all()
371-
return [self._to_channel_model(channel, db=db) for channel in channels]
377+
channel_ids = [channel.id for channel in channels]
378+
grants_map = AccessGrants.get_grants_by_resources(
379+
"channel", channel_ids, db=db
380+
)
381+
return [
382+
self._to_channel_model(
383+
channel,
384+
access_grants=grants_map.get(channel.id, []),
385+
db=db,
386+
)
387+
for channel in channels
388+
]
372389

373390
def _has_permission(self, db, query, filter: dict, permission: str = "read"):
374391
return AccessGrants.has_permission_filter(
@@ -417,7 +434,16 @@ def get_channels_by_user_id(
417434
standard_channels = query.all()
418435

419436
all_channels = membership_channels + standard_channels
420-
return [self._to_channel_model(c, db=db) for c in all_channels]
437+
channel_ids = [c.id for c in all_channels]
438+
grants_map = AccessGrants.get_grants_by_resources(
439+
"channel", channel_ids, db=db
440+
)
441+
return [
442+
self._to_channel_model(
443+
c, access_grants=grants_map.get(c.id, []), db=db
444+
)
445+
for c in all_channels
446+
]
421447

422448
def get_dm_channel_by_user_ids(
423449
self, user_ids: list[str], db: Optional[Session] = None
@@ -724,7 +750,17 @@ def get_channels_by_file_id(
724750
)
725751
channel_ids = [cf.channel_id for cf in channel_files]
726752
channels = db.query(Channel).filter(Channel.id.in_(channel_ids)).all()
727-
return [self._to_channel_model(channel, db=db) for channel in channels]
753+
grants_map = AccessGrants.get_grants_by_resources(
754+
"channel", channel_ids, db=db
755+
)
756+
return [
757+
self._to_channel_model(
758+
channel,
759+
access_grants=grants_map.get(channel.id, []),
760+
db=db,
761+
)
762+
for channel in channels
763+
]
728764

729765
def get_channels_by_file_id_and_user_id(
730766
self, file_id: str, user_id: str, db: Optional[Session] = None

backend/open_webui/models/knowledge.py

Lines changed: 36 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -144,13 +144,18 @@ def _get_access_grants(
144144
return AccessGrants.get_grants_by_resource("knowledge", knowledge_id, db=db)
145145

146146
def _to_knowledge_model(
147-
self, knowledge: Knowledge, db: Optional[Session] = None
147+
self,
148+
knowledge: Knowledge,
149+
access_grants: Optional[list[AccessGrantModel]] = None,
150+
db: Optional[Session] = None,
148151
) -> KnowledgeModel:
149152
knowledge_data = KnowledgeModel.model_validate(knowledge).model_dump(
150153
exclude={"access_grants"}
151154
)
152-
knowledge_data["access_grants"] = self._get_access_grants(
153-
knowledge_data["id"], db=db
155+
knowledge_data["access_grants"] = (
156+
access_grants
157+
if access_grants is not None
158+
else self._get_access_grants(knowledge_data["id"], db=db)
154159
)
155160
return KnowledgeModel.model_validate(knowledge_data)
156161

@@ -192,17 +197,25 @@ def get_knowledge_bases(
192197
db.query(Knowledge).order_by(Knowledge.updated_at.desc()).all()
193198
)
194199
user_ids = list(set(knowledge.user_id for knowledge in all_knowledge))
200+
knowledge_ids = [knowledge.id for knowledge in all_knowledge]
195201

196202
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
197203
users_dict = {user.id: user for user in users}
204+
grants_map = AccessGrants.get_grants_by_resources(
205+
"knowledge", knowledge_ids, db=db
206+
)
198207

199208
knowledge_bases = []
200209
for knowledge in all_knowledge:
201210
user = users_dict.get(knowledge.user_id)
202211
knowledge_bases.append(
203212
KnowledgeUserModel.model_validate(
204213
{
205-
**self._to_knowledge_model(knowledge, db=db).model_dump(),
214+
**self._to_knowledge_model(
215+
knowledge,
216+
access_grants=grants_map.get(knowledge.id, []),
217+
db=db,
218+
).model_dump(),
206219
"user": user.model_dump() if user else None,
207220
}
208221
)
@@ -261,13 +274,22 @@ def search_knowledge_bases(
261274

262275
items = query.all()
263276

277+
knowledge_ids = [kb.id for kb, _ in items]
278+
grants_map = AccessGrants.get_grants_by_resources(
279+
"knowledge", knowledge_ids, db=db
280+
)
281+
264282
knowledge_bases = []
265283
for knowledge_base, user in items:
266284
knowledge_bases.append(
267285
KnowledgeUserModel.model_validate(
268286
{
269287
**self._to_knowledge_model(
270-
knowledge_base, db=db
288+
knowledge_base,
289+
access_grants=grants_map.get(
290+
knowledge_base.id, []
291+
),
292+
db=db,
271293
).model_dump(),
272294
"user": (
273295
UserModel.model_validate(user).model_dump()
@@ -440,8 +462,16 @@ def get_knowledges_by_file_id(
440462
.filter(KnowledgeFile.file_id == file_id)
441463
.all()
442464
)
465+
knowledge_ids = [k.id for k in knowledges]
466+
grants_map = AccessGrants.get_grants_by_resources(
467+
"knowledge", knowledge_ids, db=db
468+
)
443469
return [
444-
self._to_knowledge_model(knowledge, db=db)
470+
self._to_knowledge_model(
471+
knowledge,
472+
access_grants=grants_map.get(knowledge.id, []),
473+
db=db,
474+
)
445475
for knowledge in knowledges
446476
]
447477
except Exception:

backend/open_webui/models/models.py

Lines changed: 65 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -144,11 +144,20 @@ def _get_access_grants(
144144
) -> list[AccessGrantModel]:
145145
return AccessGrants.get_grants_by_resource("model", model_id, db=db)
146146

147-
def _to_model_model(self, model: Model, db: Optional[Session] = None) -> ModelModel:
147+
def _to_model_model(
148+
self,
149+
model: Model,
150+
access_grants: Optional[list[AccessGrantModel]] = None,
151+
db: Optional[Session] = None,
152+
) -> ModelModel:
148153
model_data = ModelModel.model_validate(model).model_dump(
149154
exclude={"access_grants"}
150155
)
151-
model_data["access_grants"] = self._get_access_grants(model_data["id"], db=db)
156+
model_data["access_grants"] = (
157+
access_grants
158+
if access_grants is not None
159+
else self._get_access_grants(model_data["id"], db=db)
160+
)
152161
return ModelModel.model_validate(model_data)
153162

154163
def insert_new_model(
@@ -181,26 +190,38 @@ def insert_new_model(
181190

182191
def get_all_models(self, db: Optional[Session] = None) -> list[ModelModel]:
183192
with get_db_context(db) as db:
193+
all_models = db.query(Model).all()
194+
model_ids = [model.id for model in all_models]
195+
grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db)
184196
return [
185-
self._to_model_model(model, db=db) for model in db.query(Model).all()
197+
self._to_model_model(
198+
model, access_grants=grants_map.get(model.id, []), db=db
199+
)
200+
for model in all_models
186201
]
187202

188203
def get_models(self, db: Optional[Session] = None) -> list[ModelUserResponse]:
189204
with get_db_context(db) as db:
190205
all_models = db.query(Model).filter(Model.base_model_id != None).all()
191206

192207
user_ids = list(set(model.user_id for model in all_models))
208+
model_ids = [model.id for model in all_models]
193209

194210
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
195211
users_dict = {user.id: user for user in users}
212+
grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db)
196213

197214
models = []
198215
for model in all_models:
199216
user = users_dict.get(model.user_id)
200217
models.append(
201218
ModelUserResponse.model_validate(
202219
{
203-
**self._to_model_model(model, db=db).model_dump(),
220+
**self._to_model_model(
221+
model,
222+
access_grants=grants_map.get(model.id, []),
223+
db=db,
224+
).model_dump(),
204225
"user": user.model_dump() if user else None,
205226
}
206227
)
@@ -209,9 +230,16 @@ def get_models(self, db: Optional[Session] = None) -> list[ModelUserResponse]:
209230

210231
def get_base_models(self, db: Optional[Session] = None) -> list[ModelModel]:
211232
with get_db_context(db) as db:
233+
all_models = (
234+
db.query(Model).filter(Model.base_model_id == None).all()
235+
)
236+
model_ids = [model.id for model in all_models]
237+
grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db)
212238
return [
213-
self._to_model_model(model, db=db)
214-
for model in db.query(Model).filter(Model.base_model_id == None).all()
239+
self._to_model_model(
240+
model, access_grants=grants_map.get(model.id, []), db=db
241+
)
242+
for model in all_models
215243
]
216244

217245
def get_models_by_user_id(
@@ -325,11 +353,18 @@ def search_models(
325353

326354
items = query.all()
327355

356+
model_ids = [model.id for model, _ in items]
357+
grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db)
358+
328359
models = []
329360
for model, user in items:
330361
models.append(
331362
ModelUserResponse(
332-
**self._to_model_model(model, db=db).model_dump(),
363+
**self._to_model_model(
364+
model,
365+
access_grants=grants_map.get(model.id, []),
366+
db=db,
367+
).model_dump(),
333368
user=(
334369
UserResponse(**UserModel.model_validate(user).model_dump())
335370
if user
@@ -356,7 +391,18 @@ def get_models_by_ids(
356391
try:
357392
with get_db_context(db) as db:
358393
models = db.query(Model).filter(Model.id.in_(ids)).all()
359-
return [self._to_model_model(model, db=db) for model in models]
394+
model_ids = [model.id for model in models]
395+
grants_map = AccessGrants.get_grants_by_resources(
396+
"model", model_ids, db=db
397+
)
398+
return [
399+
self._to_model_model(
400+
model,
401+
access_grants=grants_map.get(model.id, []),
402+
db=db,
403+
)
404+
for model in models
405+
]
360406
except Exception:
361407
return []
362408

@@ -465,9 +511,18 @@ def sync_models(
465511

466512
db.commit()
467513

514+
all_models = db.query(Model).all()
515+
model_ids = [model.id for model in all_models]
516+
grants_map = AccessGrants.get_grants_by_resources(
517+
"model", model_ids, db=db
518+
)
468519
return [
469-
self._to_model_model(model, db=db)
470-
for model in db.query(Model).all()
520+
self._to_model_model(
521+
model,
522+
access_grants=grants_map.get(model.id, []),
523+
db=db,
524+
)
525+
for model in all_models
471526
]
472527
except Exception as e:
473528
log.exception(f"Error syncing models for user {user_id}: {e}")

0 commit comments

Comments
 (0)