refac/enh: db session sharing
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user