refac/enh: db session sharing
This commit is contained in:
@@ -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