refac/enh: db session sharing
This commit is contained in:
@@ -28,6 +28,7 @@ def fill_missing_permissions(
|
||||
def get_permissions(
|
||||
user_id: str,
|
||||
default_permissions: Dict[str, Any],
|
||||
db: Optional[Any] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Get all permissions for a user by combining the permissions of all groups the user is a member of.
|
||||
@@ -53,7 +54,7 @@ def get_permissions(
|
||||
) # Use the most permissive value (True > False)
|
||||
return permissions
|
||||
|
||||
user_groups = Groups.get_groups_by_member_id(user_id)
|
||||
user_groups = Groups.get_groups_by_member_id(user_id, db=db)
|
||||
|
||||
# Deep copy default permissions to avoid modifying the original dict
|
||||
permissions = json.loads(json.dumps(default_permissions))
|
||||
@@ -72,6 +73,7 @@ def has_permission(
|
||||
user_id: str,
|
||||
permission_key: str,
|
||||
default_permissions: Dict[str, Any] = {},
|
||||
db: Optional[Any] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if a user has a specific permission by checking the group permissions
|
||||
@@ -92,7 +94,7 @@ def has_permission(
|
||||
permission_hierarchy = permission_key.split(".")
|
||||
|
||||
# Retrieve user group permissions
|
||||
user_groups = Groups.get_groups_by_member_id(user_id)
|
||||
user_groups = Groups.get_groups_by_member_id(user_id, db=db)
|
||||
|
||||
for group in user_groups:
|
||||
if get_permission(group.permissions or {}, permission_hierarchy):
|
||||
@@ -127,6 +129,7 @@ def has_access(
|
||||
access_control: Optional[dict] = None,
|
||||
user_group_ids: Optional[Set[str]] = None,
|
||||
strict: bool = True,
|
||||
db: Optional[Any] = None,
|
||||
) -> bool:
|
||||
if access_control is None:
|
||||
if strict:
|
||||
@@ -135,7 +138,7 @@ def has_access(
|
||||
return True
|
||||
|
||||
if user_group_ids is None:
|
||||
user_groups = Groups.get_groups_by_member_id(user_id)
|
||||
user_groups = Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = {group.id for group in user_groups}
|
||||
|
||||
permitted_ids = get_permitted_group_and_user_ids(type, access_control)
|
||||
@@ -152,10 +155,10 @@ def has_access(
|
||||
|
||||
# Get all users with access to a resource
|
||||
def get_users_with_access(
|
||||
type: str = "write", access_control: Optional[dict] = None
|
||||
type: str = "write", access_control: Optional[dict] = None, db: Optional[Any] = None
|
||||
) -> list[UserModel]:
|
||||
if access_control is None:
|
||||
result = Users.get_users(filter={"roles": ["!pending"]})
|
||||
result = Users.get_users(filter={"roles": ["!pending"]}, db=db)
|
||||
return result.get("users", [])
|
||||
|
||||
permitted_ids = get_permitted_group_and_user_ids(type, access_control)
|
||||
@@ -167,8 +170,8 @@ def get_users_with_access(
|
||||
|
||||
user_ids_with_access = set(permitted_user_ids)
|
||||
|
||||
group_user_ids_map = Groups.get_group_user_ids_by_ids(permitted_group_ids)
|
||||
group_user_ids_map = Groups.get_group_user_ids_by_ids(permitted_group_ids, db=db)
|
||||
for user_ids in group_user_ids_map.values():
|
||||
user_ids_with_access.update(user_ids)
|
||||
|
||||
return Users.get_users_by_user_ids(list(user_ids_with_access))
|
||||
return Users.get_users_by_user_ids(list(user_ids_with_access), db=db)
|
||||
|
||||
@@ -42,6 +42,8 @@ from open_webui.env import (
|
||||
|
||||
from fastapi import BackgroundTasks, Depends, HTTPException, Request, Response, status
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import get_session
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
@@ -271,6 +273,7 @@ async def get_current_user(
|
||||
response: Response,
|
||||
background_tasks: BackgroundTasks,
|
||||
auth_token: HTTPAuthorizationCredentials = Depends(bearer_security),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
token = None
|
||||
|
||||
@@ -285,7 +288,7 @@ async def get_current_user(
|
||||
|
||||
# auth by api key
|
||||
if token.startswith("sk-"):
|
||||
user = get_current_user_by_api_key(request, token)
|
||||
user = get_current_user_by_api_key(request, token, db=db)
|
||||
|
||||
# Add user info to current span
|
||||
current_span = trace.get_current_span()
|
||||
@@ -314,7 +317,7 @@ async def get_current_user(
|
||||
detail="Invalid token",
|
||||
)
|
||||
|
||||
user = Users.get_user_by_id(data["id"])
|
||||
user = Users.get_user_by_id(data["id"], db=db)
|
||||
if user is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -364,8 +367,8 @@ async def get_current_user(
|
||||
raise e
|
||||
|
||||
|
||||
def get_current_user_by_api_key(request, api_key: str):
|
||||
user = Users.get_user_by_api_key(api_key)
|
||||
def get_current_user_by_api_key(request, api_key: str, db: Session = None):
|
||||
user = Users.get_user_by_api_key(api_key, db=db)
|
||||
|
||||
if user is None:
|
||||
raise HTTPException(
|
||||
@@ -393,7 +396,7 @@ def get_current_user_by_api_key(request, api_key: str):
|
||||
current_span.set_attribute("client.user.role", user.role)
|
||||
current_span.set_attribute("client.auth.type", "api_key")
|
||||
|
||||
Users.update_last_active_by_id(user.id)
|
||||
Users.update_last_active_by_id(user.id, db=db)
|
||||
return user
|
||||
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ log = logging.getLogger(__name__)
|
||||
def apply_default_group_assignment(
|
||||
default_group_id: str,
|
||||
user_id: str,
|
||||
db=None,
|
||||
) -> None:
|
||||
"""
|
||||
Apply default group assignment to a user if default_group_id is provided.
|
||||
@@ -17,7 +18,7 @@ def apply_default_group_assignment(
|
||||
"""
|
||||
if default_group_id:
|
||||
try:
|
||||
Groups.add_users_to_group(default_group_id, [user_id])
|
||||
Groups.add_users_to_group(default_group_id, [user_id], db=db)
|
||||
except Exception as e:
|
||||
log.error(
|
||||
f"Failed to add user {user_id} to default group {default_group_id}: {e}"
|
||||
|
||||
@@ -1336,7 +1336,7 @@ class OAuthManager:
|
||||
|
||||
return await client.authorize_redirect(request, redirect_uri, **kwargs)
|
||||
|
||||
async def handle_callback(self, request, provider, response):
|
||||
async def handle_callback(self, request, provider, response, db=None):
|
||||
if provider not in OAUTH_PROVIDERS:
|
||||
raise HTTPException(404)
|
||||
|
||||
@@ -1461,20 +1461,20 @@ class OAuthManager:
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
|
||||
|
||||
# Check if the user exists
|
||||
user = Users.get_user_by_oauth_sub(provider, sub)
|
||||
user = Users.get_user_by_oauth_sub(provider, sub, db=db)
|
||||
if not user:
|
||||
# If the user does not exist, check if merging is enabled
|
||||
if auth_manager_config.OAUTH_MERGE_ACCOUNTS_BY_EMAIL:
|
||||
# Check if the user exists by email
|
||||
user = Users.get_user_by_email(email)
|
||||
user = Users.get_user_by_email(email, db=db)
|
||||
if user:
|
||||
# Update the user with the new oauth sub
|
||||
Users.update_user_oauth_by_id(user.id, provider, sub)
|
||||
Users.update_user_oauth_by_id(user.id, provider, sub, db=db)
|
||||
|
||||
if user:
|
||||
determined_role = self.get_user_role(user, user_data)
|
||||
if user.role != determined_role:
|
||||
Users.update_user_role_by_id(user.id, determined_role)
|
||||
Users.update_user_role_by_id(user.id, determined_role, db=db)
|
||||
# Update the user object in memory as well,
|
||||
# to avoid problems with the ENABLE_OAUTH_GROUP_MANAGEMENT check below
|
||||
user.role = determined_role
|
||||
@@ -1491,14 +1491,14 @@ class OAuthManager:
|
||||
)
|
||||
if processed_picture_url != user.profile_image_url:
|
||||
Users.update_user_profile_image_url_by_id(
|
||||
user.id, processed_picture_url
|
||||
user.id, processed_picture_url, db=db
|
||||
)
|
||||
log.debug(f"Updated profile picture for user {user.email}")
|
||||
else:
|
||||
# If the user does not exist, check if signups are enabled
|
||||
if auth_manager_config.ENABLE_OAUTH_SIGNUP:
|
||||
# Check if an existing user with the same email already exists
|
||||
existing_user = Users.get_user_by_email(email)
|
||||
existing_user = Users.get_user_by_email(email, db=db)
|
||||
if existing_user:
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN)
|
||||
|
||||
@@ -1529,6 +1529,7 @@ class OAuthManager:
|
||||
profile_image_url=picture_url,
|
||||
role=self.get_user_role(None, user_data),
|
||||
oauth=oauth_data,
|
||||
db=db,
|
||||
)
|
||||
|
||||
if auth_manager_config.WEBHOOK_URL:
|
||||
@@ -1544,8 +1545,7 @@ class OAuthManager:
|
||||
)
|
||||
|
||||
apply_default_group_assignment(
|
||||
request.app.state.config.DEFAULT_GROUP_ID,
|
||||
user.id,
|
||||
request.app.state.config.DEFAULT_GROUP_ID, user.id, db=db
|
||||
)
|
||||
|
||||
else:
|
||||
@@ -1616,15 +1616,16 @@ class OAuthManager:
|
||||
token["expires_at"] = datetime.now().timestamp() + token["expires_in"]
|
||||
|
||||
# Clean up any existing sessions for this user/provider first
|
||||
sessions = OAuthSessions.get_sessions_by_user_id(user.id)
|
||||
sessions = OAuthSessions.get_sessions_by_user_id(user.id, db=db)
|
||||
for session in sessions:
|
||||
if session.provider == provider:
|
||||
OAuthSessions.delete_session_by_id(session.id)
|
||||
OAuthSessions.delete_session_by_id(session.id, db=db)
|
||||
|
||||
session = OAuthSessions.create_session(
|
||||
user_id=user.id,
|
||||
provider=provider,
|
||||
token=token,
|
||||
db=db,
|
||||
)
|
||||
|
||||
response.set_cookie(
|
||||
|
||||
Reference in New Issue
Block a user