refac/enh: db session sharing

This commit is contained in:
Timothy Jaeryang Baek
2025-12-29 00:21:18 +04:00
parent 6dd0f99b90
commit b1d0f00d8c
23 changed files with 1173 additions and 663 deletions
+83 -46
View File
@@ -1,5 +1,6 @@
import logging
from typing import Optional
from sqlalchemy.orm import Session
import base64
import io
@@ -29,6 +30,7 @@ from open_webui.models.users import (
from open_webui.constants import ERROR_MESSAGES
from open_webui.env import STATIC_DIR
from open_webui.internal.db import get_session
from open_webui.utils.auth import (
@@ -60,6 +62,7 @@ async def get_users(
direction: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_admin_user),
db: Session = Depends(get_session),
):
limit = PAGE_ITEM_COUNT
@@ -74,7 +77,9 @@ async def get_users(
if direction:
filter["direction"] = direction
result = Users.get_users(filter=filter, skip=skip, limit=limit)
filter["direction"] = direction
result = Users.get_users(filter=filter, skip=skip, limit=limit, db=db)
users = result["users"]
total = result["total"]
@@ -85,7 +90,8 @@ async def get_users(
**{
**user.model_dump(),
"group_ids": [
group.id for group in Groups.get_groups_by_member_id(user.id)
group.id
for group in Groups.get_groups_by_member_id(user.id, db=db)
],
}
)
@@ -98,8 +104,9 @@ async def get_users(
@router.get("/all", response_model=UserInfoListResponse)
async def get_all_users(
user=Depends(get_admin_user),
db: Session = Depends(get_session),
):
return Users.get_users()
return Users.get_users(db=db)
@router.get("/search", response_model=UserInfoListResponse)
@@ -109,16 +116,13 @@ async def search_users(
direction: Optional[str] = None,
page: Optional[int] = 1,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
limit = PAGE_ITEM_COUNT
page = max(1, page)
skip = (page - 1) * limit
filter = {}
if query:
filter["query"] = query
filter = {}
if query:
filter["query"] = query
@@ -127,7 +131,7 @@ async def search_users(
if direction:
filter["direction"] = direction
return Users.get_users(filter=filter, skip=skip, limit=limit)
return Users.get_users(filter=filter, skip=skip, limit=limit, db=db)
############################
@@ -136,8 +140,10 @@ async def search_users(
@router.get("/groups")
async def get_user_groups(user=Depends(get_verified_user)):
return Groups.get_groups_by_member_id(user.id)
async def get_user_groups(
user=Depends(get_verified_user), db: Session = Depends(get_session)
):
return Groups.get_groups_by_member_id(user.id, db=db)
############################
@@ -146,9 +152,13 @@ async def get_user_groups(user=Depends(get_verified_user)):
@router.get("/permissions")
async def get_user_permissisions(request: Request, user=Depends(get_verified_user)):
async def get_user_permissisions(
request: Request,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
user_permissions = get_permissions(
user.id, request.app.state.config.USER_PERMISSIONS
user.id, request.app.state.config.USER_PERMISSIONS, db=db
)
return user_permissions
@@ -256,8 +266,10 @@ async def update_default_user_permissions(
@router.get("/user/settings", response_model=Optional[UserSettings])
async def get_user_settings_by_session_user(user=Depends(get_verified_user)):
user = Users.get_user_by_id(user.id)
async def get_user_settings_by_session_user(
user=Depends(get_verified_user), db: Session = Depends(get_session)
):
user = Users.get_user_by_id(user.id, db=db)
if user:
return user.settings
else:
@@ -274,7 +286,10 @@ async def get_user_settings_by_session_user(user=Depends(get_verified_user)):
@router.post("/user/settings/update", response_model=UserSettings)
async def update_user_settings_by_session_user(
request: Request, form_data: UserSettings, user=Depends(get_verified_user)
request: Request,
form_data: UserSettings,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
updated_user_settings = form_data.model_dump()
if (
@@ -289,7 +304,7 @@ async def update_user_settings_by_session_user(
# If the user is not an admin and does not have permission to use tool servers, remove the key
updated_user_settings["ui"].pop("toolServers", None)
user = Users.update_user_settings_by_id(user.id, updated_user_settings)
user = Users.update_user_settings_by_id(user.id, updated_user_settings, db=db)
if user:
return user.settings
else:
@@ -305,8 +320,10 @@ async def update_user_settings_by_session_user(
@router.get("/user/status")
async def get_user_status_by_session_user(user=Depends(get_verified_user)):
user = Users.get_user_by_id(user.id)
async def get_user_status_by_session_user(
user=Depends(get_verified_user), db: Session = Depends(get_session)
):
user = Users.get_user_by_id(user.id, db=db)
if user:
return user
else:
@@ -323,11 +340,13 @@ async def get_user_status_by_session_user(user=Depends(get_verified_user)):
@router.post("/user/status/update")
async def update_user_status_by_session_user(
form_data: UserStatus, user=Depends(get_verified_user)
form_data: UserStatus,
user=Depends(get_verified_user),
db: Session = Depends(get_session),
):
user = Users.get_user_by_id(user.id)
user = Users.get_user_by_id(user.id, db=db)
if user:
user = Users.update_user_status_by_id(user.id, form_data)
user = Users.update_user_status_by_id(user.id, form_data, db=db)
return user
else:
raise HTTPException(
@@ -342,8 +361,10 @@ async def update_user_status_by_session_user(
@router.get("/user/info", response_model=Optional[dict])
async def get_user_info_by_session_user(user=Depends(get_verified_user)):
user = Users.get_user_by_id(user.id)
async def get_user_info_by_session_user(
user=Depends(get_verified_user), db: Session = Depends(get_session)
):
user = Users.get_user_by_id(user.id, db=db)
if user:
return user.info
else:
@@ -360,14 +381,16 @@ async def get_user_info_by_session_user(user=Depends(get_verified_user)):
@router.post("/user/info/update", response_model=Optional[dict])
async def update_user_info_by_session_user(
form_data: dict, user=Depends(get_verified_user)
form_data: dict, user=Depends(get_verified_user), db: Session = Depends(get_session)
):
user = Users.get_user_by_id(user.id)
user = Users.get_user_by_id(user.id, db=db)
if user:
if user.info is None:
user.info = {}
user = Users.update_user_by_id(user.id, {"info": {**user.info, **form_data}})
user = Users.update_user_by_id(
user.id, {"info": {**user.info, **form_data}}, db=db
)
if user:
return user.info
else:
@@ -397,7 +420,9 @@ class UserActiveResponse(UserStatus):
@router.get("/{user_id}", response_model=UserActiveResponse)
async def get_user_by_id(user_id: str, user=Depends(get_verified_user)):
async def get_user_by_id(
user_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
):
# Check if user_id is a shared chat
# If it is, get the user_id from the chat
if user_id.startswith("shared-"):
@@ -411,14 +436,14 @@ async def get_user_by_id(user_id: str, user=Depends(get_verified_user)):
detail=ERROR_MESSAGES.USER_NOT_FOUND,
)
user = Users.get_user_by_id(user_id)
user = Users.get_user_by_id(user_id, db=db)
if user:
groups = Groups.get_groups_by_member_id(user_id)
groups = Groups.get_groups_by_member_id(user_id, db=db)
return UserActiveResponse(
**{
**user.model_dump(),
"groups": [{"id": group.id, "name": group.name} for group in groups],
"is_active": Users.is_user_active(user_id),
"is_active": Users.is_user_active(user_id, db=db),
}
)
else:
@@ -429,8 +454,10 @@ async def get_user_by_id(user_id: str, user=Depends(get_verified_user)):
@router.get("/{user_id}/oauth/sessions")
async def get_user_oauth_sessions_by_id(user_id: str, user=Depends(get_admin_user)):
sessions = OAuthSessions.get_sessions_by_user_id(user_id)
async def get_user_oauth_sessions_by_id(
user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)
):
sessions = OAuthSessions.get_sessions_by_user_id(user_id, db=db)
if sessions and len(sessions) > 0:
return sessions
else:
@@ -446,8 +473,10 @@ async def get_user_oauth_sessions_by_id(user_id: str, user=Depends(get_admin_use
@router.get("/{user_id}/profile/image")
async def get_user_profile_image_by_id(user_id: str, user=Depends(get_verified_user)):
user = Users.get_user_by_id(user_id)
async def get_user_profile_image_by_id(
user_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
):
user = Users.get_user_by_id(user_id, db=db)
if user:
if user.profile_image_url:
# check if it's url or base64
@@ -484,9 +513,11 @@ async def get_user_profile_image_by_id(user_id: str, user=Depends(get_verified_u
@router.get("/{user_id}/active", response_model=dict)
async def get_user_active_status_by_id(user_id: str, user=Depends(get_verified_user)):
async def get_user_active_status_by_id(
user_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
):
return {
"active": Users.is_user_active(user_id),
"active": Users.is_user_active(user_id, db=db),
}
@@ -500,10 +531,11 @@ async def update_user_by_id(
user_id: str,
form_data: UserUpdateForm,
session_user=Depends(get_admin_user),
db: Session = Depends(get_session),
):
# Prevent modification of the primary admin user by other admins
try:
first_user = Users.get_first_user()
first_user = Users.get_first_user(db=db)
if first_user:
if user_id == first_user.id:
if session_user.id != user_id:
@@ -527,11 +559,11 @@ async def update_user_by_id(
detail="Could not verify primary admin status.",
)
user = Users.get_user_by_id(user_id)
user = Users.get_user_by_id(user_id, db=db)
if user:
if form_data.email.lower() != user.email:
email_user = Users.get_user_by_email(form_data.email.lower())
email_user = Users.get_user_by_email(form_data.email.lower(), db=db)
if email_user:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
@@ -545,9 +577,9 @@ async def update_user_by_id(
raise HTTPException(400, detail=str(e))
hashed = get_password_hash(form_data.password)
Auths.update_user_password_by_id(user_id, hashed)
Auths.update_user_password_by_id(user_id, hashed, db=db)
Auths.update_email_by_id(user_id, form_data.email.lower())
Auths.update_email_by_id(user_id, form_data.email.lower(), db=db)
updated_user = Users.update_user_by_id(
user_id,
{
@@ -556,6 +588,7 @@ async def update_user_by_id(
"email": form_data.email.lower(),
"profile_image_url": form_data.profile_image_url,
},
db=db,
)
if updated_user:
@@ -578,10 +611,12 @@ async def update_user_by_id(
@router.delete("/{user_id}", response_model=bool)
async def delete_user_by_id(user_id: str, user=Depends(get_admin_user)):
async def delete_user_by_id(
user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)
):
# Prevent deletion of the primary admin user
try:
first_user = Users.get_first_user()
first_user = Users.get_first_user(db=db)
if first_user and user_id == first_user.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
@@ -595,7 +630,7 @@ async def delete_user_by_id(user_id: str, user=Depends(get_admin_user)):
)
if user.id != user_id:
result = Auths.delete_auth_by_id(user_id)
result = Auths.delete_auth_by_id(user_id, db=db)
if result:
return True
@@ -618,5 +653,7 @@ async def delete_user_by_id(user_id: str, user=Depends(get_admin_user)):
@router.get("/{user_id}/groups")
async def get_user_groups_by_id(user_id: str, user=Depends(get_admin_user)):
return Groups.get_groups_by_member_id(user_id)
async def get_user_groups_by_id(
user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)
):
return Groups.get_groups_by_member_id(user_id, db=db)