refac/enh: db session sharing
This commit is contained in:
@@ -16,6 +16,9 @@ from open_webui.config import CACHE_DIR
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
|
||||
from open_webui.internal.db import get_session
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
|
||||
|
||||
@@ -29,7 +32,11 @@ router = APIRouter()
|
||||
|
||||
|
||||
@router.get("/", response_model=list[GroupResponse])
|
||||
async def get_groups(share: Optional[bool] = None, user=Depends(get_verified_user)):
|
||||
async def get_groups(
|
||||
share: Optional[bool] = None,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
|
||||
filter = {}
|
||||
if user.role != "admin":
|
||||
@@ -38,7 +45,7 @@ async def get_groups(share: Optional[bool] = None, user=Depends(get_verified_use
|
||||
if share is not None:
|
||||
filter["share"] = share
|
||||
|
||||
groups = Groups.get_groups(filter=filter)
|
||||
groups = Groups.get_groups(filter=filter, db=db)
|
||||
|
||||
return groups
|
||||
|
||||
@@ -49,13 +56,17 @@ async def get_groups(share: Optional[bool] = None, user=Depends(get_verified_use
|
||||
|
||||
|
||||
@router.post("/create", response_model=Optional[GroupResponse])
|
||||
async def create_new_group(form_data: GroupForm, user=Depends(get_admin_user)):
|
||||
async def create_new_group(
|
||||
form_data: GroupForm,
|
||||
user=Depends(get_admin_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
try:
|
||||
group = Groups.insert_new_group(user.id, form_data)
|
||||
group = Groups.insert_new_group(user.id, form_data, db=db)
|
||||
if group:
|
||||
return GroupResponse(
|
||||
**group.model_dump(),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id, db=db),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -76,12 +87,14 @@ async def create_new_group(form_data: GroupForm, user=Depends(get_admin_user)):
|
||||
|
||||
|
||||
@router.get("/id/{id}", response_model=Optional[GroupResponse])
|
||||
async def get_group_by_id(id: str, user=Depends(get_admin_user)):
|
||||
group = Groups.get_group_by_id(id)
|
||||
async def get_group_by_id(
|
||||
id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
group = Groups.get_group_by_id(id, db=db)
|
||||
if group:
|
||||
return GroupResponse(
|
||||
**group.model_dump(),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id, db=db),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -101,13 +114,15 @@ class GroupExportResponse(GroupResponse):
|
||||
|
||||
|
||||
@router.get("/id/{id}/export", response_model=Optional[GroupExportResponse])
|
||||
async def export_group_by_id(id: str, user=Depends(get_admin_user)):
|
||||
group = Groups.get_group_by_id(id)
|
||||
async def export_group_by_id(
|
||||
id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
group = Groups.get_group_by_id(id, db=db)
|
||||
if group:
|
||||
return GroupExportResponse(
|
||||
**group.model_dump(),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id),
|
||||
user_ids=Groups.get_group_user_ids_by_id(group.id),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id, db=db),
|
||||
user_ids=Groups.get_group_user_ids_by_id(group.id, db=db),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -122,9 +137,11 @@ async def export_group_by_id(id: str, user=Depends(get_admin_user)):
|
||||
|
||||
|
||||
@router.post("/id/{id}/users", response_model=list[UserInfoResponse])
|
||||
async def get_users_in_group(id: str, user=Depends(get_admin_user)):
|
||||
async def get_users_in_group(
|
||||
id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
try:
|
||||
users = Users.get_users_by_group_id(id)
|
||||
users = Users.get_users_by_group_id(id, db=db)
|
||||
return users
|
||||
except Exception as e:
|
||||
log.exception(f"Error adding users to group {id}: {e}")
|
||||
@@ -141,14 +158,17 @@ async def get_users_in_group(id: str, user=Depends(get_admin_user)):
|
||||
|
||||
@router.post("/id/{id}/update", response_model=Optional[GroupResponse])
|
||||
async def update_group_by_id(
|
||||
id: str, form_data: GroupUpdateForm, user=Depends(get_admin_user)
|
||||
id: str,
|
||||
form_data: GroupUpdateForm,
|
||||
user=Depends(get_admin_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
try:
|
||||
group = Groups.update_group_by_id(id, form_data)
|
||||
group = Groups.update_group_by_id(id, form_data, db=db)
|
||||
if group:
|
||||
return GroupResponse(
|
||||
**group.model_dump(),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id, db=db),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -170,17 +190,20 @@ async def update_group_by_id(
|
||||
|
||||
@router.post("/id/{id}/users/add", response_model=Optional[GroupResponse])
|
||||
async def add_user_to_group(
|
||||
id: str, form_data: UserIdsForm, user=Depends(get_admin_user)
|
||||
id: str,
|
||||
form_data: UserIdsForm,
|
||||
user=Depends(get_admin_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
try:
|
||||
if form_data.user_ids:
|
||||
form_data.user_ids = Users.get_valid_user_ids(form_data.user_ids)
|
||||
form_data.user_ids = Users.get_valid_user_ids(form_data.user_ids, db=db)
|
||||
|
||||
group = Groups.add_users_to_group(id, form_data.user_ids)
|
||||
group = Groups.add_users_to_group(id, form_data.user_ids, db=db)
|
||||
if group:
|
||||
return GroupResponse(
|
||||
**group.model_dump(),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id, db=db),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -197,14 +220,17 @@ async def add_user_to_group(
|
||||
|
||||
@router.post("/id/{id}/users/remove", response_model=Optional[GroupResponse])
|
||||
async def remove_users_from_group(
|
||||
id: str, form_data: UserIdsForm, user=Depends(get_admin_user)
|
||||
id: str,
|
||||
form_data: UserIdsForm,
|
||||
user=Depends(get_admin_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
try:
|
||||
group = Groups.remove_users_from_group(id, form_data.user_ids)
|
||||
group = Groups.remove_users_from_group(id, form_data.user_ids, db=db)
|
||||
if group:
|
||||
return GroupResponse(
|
||||
**group.model_dump(),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id),
|
||||
member_count=Groups.get_group_member_count_by_id(group.id, db=db),
|
||||
)
|
||||
else:
|
||||
raise HTTPException(
|
||||
@@ -225,9 +251,11 @@ async def remove_users_from_group(
|
||||
|
||||
|
||||
@router.delete("/id/{id}/delete", response_model=bool)
|
||||
async def delete_group_by_id(id: str, user=Depends(get_admin_user)):
|
||||
async def delete_group_by_id(
|
||||
id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
try:
|
||||
result = Groups.delete_group_by_id(id)
|
||||
result = Groups.delete_group_by_id(id, db=db)
|
||||
if result:
|
||||
return result
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user