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