refac/enh: db session sharing

This commit is contained in:
Timothy Jaeryang Baek
2025-12-28 22:00:44 +04:00
parent d4de26bd05
commit 2041ab483e
20 changed files with 600 additions and 562 deletions
+42 -41
View File
@@ -4,7 +4,8 @@ import time
from typing import Optional
import uuid
from open_webui.internal.db import Base, get_db
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from open_webui.models.files import FileMetadataResponse
@@ -120,9 +121,9 @@ class GroupListResponse(BaseModel):
class GroupTable:
def insert_new_group(
self, user_id: str, form_data: GroupForm
self, user_id: str, form_data: GroupForm, db: Optional[Session] = None
) -> Optional[GroupModel]:
with get_db() as db:
with get_db_context(db) as db:
group = GroupModel(
**{
**form_data.model_dump(exclude_none=True),
@@ -146,13 +147,13 @@ class GroupTable:
except Exception:
return None
def get_all_groups(self) -> list[GroupModel]:
with get_db() as db:
def get_all_groups(self, db: Optional[Session] = None) -> list[GroupModel]:
with get_db_context(db) as db:
groups = db.query(Group).order_by(Group.updated_at.desc()).all()
return [GroupModel.model_validate(group) for group in groups]
def get_groups(self, filter) -> list[GroupResponse]:
with get_db() as db:
def get_groups(self, filter, db: Optional[Session] = None) -> list[GroupResponse]:
with get_db_context(db) as db:
query = db.query(Group)
if filter:
@@ -184,16 +185,16 @@ class GroupTable:
GroupResponse.model_validate(
{
**GroupModel.model_validate(group).model_dump(),
"member_count": self.get_group_member_count_by_id(group.id),
"member_count": self.get_group_member_count_by_id(group.id, db=db),
}
)
for group in groups
]
def search_groups(
self, filter: Optional[dict] = None, skip: int = 0, limit: int = 30
self, filter: Optional[dict] = None, skip: int = 0, limit: int = 30, db: Optional[Session] = None
) -> GroupListResponse:
with get_db() as db:
with get_db_context(db) as db:
query = db.query(Group)
if filter:
@@ -220,15 +221,15 @@ class GroupTable:
"items": [
GroupResponse.model_validate(
**GroupModel.model_validate(group).model_dump(),
member_count=self.get_group_member_count_by_id(group.id),
member_count=self.get_group_member_count_by_id(group.id, db=db),
)
for group in groups
],
"total": total,
}
def get_groups_by_member_id(self, user_id: str) -> list[GroupModel]:
with get_db() as db:
def get_groups_by_member_id(self, user_id: str, db: Optional[Session] = None) -> list[GroupModel]:
with get_db_context(db) as db:
return [
GroupModel.model_validate(group)
for group in db.query(Group)
@@ -238,16 +239,16 @@ class GroupTable:
.all()
]
def get_group_by_id(self, id: str) -> Optional[GroupModel]:
def get_group_by_id(self, id: str, db: Optional[Session] = None) -> Optional[GroupModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
group = db.query(Group).filter_by(id=id).first()
return GroupModel.model_validate(group) if group else None
except Exception:
return None
def get_group_user_ids_by_id(self, id: str) -> Optional[list[str]]:
with get_db() as db:
def get_group_user_ids_by_id(self, id: str, db: Optional[Session] = None) -> Optional[list[str]]:
with get_db_context(db) as db:
members = (
db.query(GroupMember.user_id).filter(GroupMember.group_id == id).all()
)
@@ -257,8 +258,8 @@ class GroupTable:
return [m[0] for m in members]
def get_group_user_ids_by_ids(self, group_ids: list[str]) -> dict[str, list[str]]:
with get_db() as db:
def get_group_user_ids_by_ids(self, group_ids: list[str], db: Optional[Session] = None) -> dict[str, list[str]]:
with get_db_context(db) as db:
members = (
db.query(GroupMember.group_id, GroupMember.user_id)
.filter(GroupMember.group_id.in_(group_ids))
@@ -274,8 +275,8 @@ class GroupTable:
return group_user_ids
def set_group_user_ids_by_id(self, group_id: str, user_ids: list[str]) -> None:
with get_db() as db:
def set_group_user_ids_by_id(self, group_id: str, user_ids: list[str], db: Optional[Session] = None) -> None:
with get_db_context(db) as db:
# Delete existing members
db.query(GroupMember).filter(GroupMember.group_id == group_id).delete()
@@ -295,8 +296,8 @@ class GroupTable:
db.add_all(new_members)
db.commit()
def get_group_member_count_by_id(self, id: str) -> int:
with get_db() as db:
def get_group_member_count_by_id(self, id: str, db: Optional[Session] = None) -> int:
with get_db_context(db) as db:
count = (
db.query(func.count(GroupMember.user_id))
.filter(GroupMember.group_id == id)
@@ -305,10 +306,10 @@ class GroupTable:
return count if count else 0
def update_group_by_id(
self, id: str, form_data: GroupUpdateForm, overwrite: bool = False
self, id: str, form_data: GroupUpdateForm, overwrite: bool = False, db: Optional[Session] = None
) -> Optional[GroupModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
db.query(Group).filter_by(id=id).update(
{
**form_data.model_dump(exclude_none=True),
@@ -316,22 +317,22 @@ class GroupTable:
}
)
db.commit()
return self.get_group_by_id(id=id)
return self.get_group_by_id(id=id, db=db)
except Exception as e:
log.exception(e)
return None
def delete_group_by_id(self, id: str) -> bool:
def delete_group_by_id(self, id: str, db: Optional[Session] = None) -> bool:
try:
with get_db() as db:
with get_db_context(db) as db:
db.query(Group).filter_by(id=id).delete()
db.commit()
return True
except Exception:
return False
def delete_all_groups(self) -> bool:
with get_db() as db:
def delete_all_groups(self, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
try:
db.query(Group).delete()
db.commit()
@@ -340,8 +341,8 @@ class GroupTable:
except Exception:
return False
def remove_user_from_all_groups(self, user_id: str) -> bool:
with get_db() as db:
def remove_user_from_all_groups(self, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
try:
# Find all groups the user belongs to
groups = (
@@ -369,16 +370,16 @@ class GroupTable:
return False
def create_groups_by_group_names(
self, user_id: str, group_names: list[str]
self, user_id: str, group_names: list[str], db: Optional[Session] = None
) -> list[GroupModel]:
# check for existing groups
existing_groups = self.get_all_groups()
existing_groups = self.get_all_groups(db=db)
existing_group_names = {group.name for group in existing_groups}
new_groups = []
with get_db() as db:
with get_db_context(db) as db:
for group_name in group_names:
if group_name not in existing_group_names:
new_group = GroupModel(
@@ -400,8 +401,8 @@ class GroupTable:
continue
return new_groups
def sync_groups_by_group_names(self, user_id: str, group_names: list[str]) -> bool:
with get_db() as db:
def sync_groups_by_group_names(self, user_id: str, group_names: list[str], db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
try:
now = int(time.time())
@@ -461,10 +462,10 @@ class GroupTable:
return False
def add_users_to_group(
self, id: str, user_ids: Optional[list[str]] = None
self, id: str, user_ids: Optional[list[str]] = None, db: Optional[Session] = None
) -> Optional[GroupModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
group = db.query(Group).filter_by(id=id).first()
if not group:
return None
@@ -499,10 +500,10 @@ class GroupTable:
return None
def remove_users_from_group(
self, id: str, user_ids: Optional[list[str]] = None
self, id: str, user_ids: Optional[list[str]] = None, db: Optional[Session] = None
) -> Optional[GroupModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
group = db.query(Group).filter_by(id=id).first()
if not group:
return None