Skip to content

Commit b67a60b

Browse files
OAuth redirect fixes
1 parent b891df9 commit b67a60b

3 files changed

Lines changed: 55 additions & 33 deletions

File tree

cps/__init__.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -225,6 +225,11 @@ def _cwa_ensure_db_session():
225225
# Failsafe: let route-level code handle specific DB errors
226226
pass
227227

228+
@app.teardown_appcontext
229+
def shutdown_session(exception=None):
230+
if calibre_db.session_factory:
231+
calibre_db.session_factory.remove()
232+
228233
# Load user from reverse proxy header early in request lifecycle
229234
# This ensures current_user resolves correctly before any code accesses user settings
230235
@app.before_request

cps/oauth.py

Lines changed: 29 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -80,33 +80,53 @@ def set(self, blueprint, token, user=None, user_id=None):
8080
u = first(_get_real_user(ref, self.anon_user)
8181
for ref in (user, self.user, blueprint.config.get("user")))
8282

83-
if self.user_required and not u and not uid:
83+
# Check session for provider_user_id (Binding flow)
84+
provider_user_id = None
85+
if self.provider_id + '_oauth_user_id' in session:
86+
provider_user_id = session[self.provider_id + '_oauth_user_id']
87+
88+
if self.user_required and not u and not uid and not provider_user_id:
8489
raise ValueError("Cannot set OAuth token without an associated user")
8590

8691
# if there was an existing model, delete it
8792
existing_query = (
8893
self.session.query(self.model)
8994
.filter_by(provider=self.provider_id)
9095
)
91-
# check for user ID
92-
has_user_id = hasattr(self.model, "user_id")
93-
if has_user_id and uid:
94-
existing_query = existing_query.filter_by(user_id=uid)
95-
# check for user (relationship property)
96-
has_user = hasattr(self.model, "user")
97-
if has_user and u:
98-
existing_query = existing_query.filter_by(user=u)
96+
97+
if provider_user_id:
98+
existing_query = existing_query.filter_by(provider_user_id=provider_user_id)
99+
else:
100+
# check for user ID
101+
has_user_id = hasattr(self.model, "user_id")
102+
if has_user_id and uid:
103+
existing_query = existing_query.filter_by(user_id=uid)
104+
# check for user (relationship property)
105+
has_user = hasattr(self.model, "user")
106+
if has_user and u:
107+
existing_query = existing_query.filter_by(user=u)
108+
99109
# queue up delete query -- won't be run until commit()
100110
existing_query.delete()
111+
101112
# create a new model for this token
102113
kwargs = {
103114
"provider": self.provider_id,
104115
"token": token,
105116
}
117+
if provider_user_id:
118+
kwargs["provider_user_id"] = provider_user_id
119+
120+
# Only set user if we have a valid user (not anonymous/None)
121+
# Re-check has_user/has_user_id since they were in else block before
122+
has_user_id = hasattr(self.model, "user_id")
123+
has_user = hasattr(self.model, "user")
124+
106125
if has_user_id and uid:
107126
kwargs["user_id"] = uid
108127
if has_user and u:
109128
kwargs["user"] = u
129+
110130
self.session.add(self.model(**kwargs))
111131
# commit to delete and add simultaneously
112132
self.session.commit()

cps/oauth_bb.py

Lines changed: 21 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -484,19 +484,29 @@ def generate_oauth_blueprints():
484484
ub.session_commit("{} Blueprint Created".format(provider))
485485

486486
oauth_ids = ub.session.query(ub.OAuthProvider).filter(ub.OAuthProvider.provider_name.in_(['github', 'google'])).all()
487+
488+
# Ensure deterministic assignment of providers regardless of DB query order
489+
github_provider = next((p for p in oauth_ids if p.provider_name == 'github'), None)
490+
google_provider = next((p for p in oauth_ids if p.provider_name == 'google'), None)
491+
492+
# Fallback if providers are missing (shouldn't happen due to creation logic above)
493+
if not github_provider or not google_provider:
494+
log.error("OAuth providers missing after creation check")
495+
return []
496+
487497
ele1 = dict(provider_name='github',
488-
id=oauth_ids[0].id,
489-
active=oauth_ids[0].active,
490-
oauth_client_id=oauth_ids[0].oauth_client_id,
498+
id=github_provider.id,
499+
active=github_provider.active,
500+
oauth_client_id=github_provider.oauth_client_id,
491501
scope=None,
492-
oauth_client_secret=oauth_ids[0].oauth_client_secret,
502+
oauth_client_secret=github_provider.oauth_client_secret,
493503
obtain_link='https://github.com/settings/developers')
494504
ele2 = dict(provider_name='google',
495-
id=oauth_ids[1].id,
496-
active=oauth_ids[1].active,
505+
id=google_provider.id,
506+
active=google_provider.active,
497507
scope=["https://www.googleapis.com/auth/userinfo.email"],
498-
oauth_client_id=oauth_ids[1].oauth_client_id,
499-
oauth_client_secret=oauth_ids[1].oauth_client_secret,
508+
oauth_client_id=google_provider.oauth_client_id,
509+
oauth_client_secret=google_provider.oauth_client_secret,
500510
obtain_link='https://console.developers.google.com/apis/credentials')
501511
oauthblueprints.append(ele1)
502512
oauthblueprints.append(ele2)
@@ -781,12 +791,9 @@ def github_login():
781791
del github.token
782792
except Exception:
783793
pass
784-
# Force clear session to prevent loops
785-
if 'github_oauth_token' in session: session.pop('github_oauth_token', None)
786-
if 'github_oauth_user_id' in session: session.pop('github_oauth_user_id', None)
787-
788794
flash(_("GitHub Oauth error: {}").format(e), category="error")
789795
log.error(e)
796+
return redirect(url_for('github.login'))
790797
return redirect(url_for('web.login'))
791798

792799

@@ -813,12 +820,9 @@ def google_login():
813820
del google.token
814821
except Exception:
815822
pass
816-
# Force clear session to prevent loops
817-
if 'google_oauth_token' in session: session.pop('google_oauth_token', None)
818-
if 'google_oauth_user_id' in session: session.pop('google_oauth_user_id', None)
819-
820823
flash(_("Google Oauth error: {}").format(e), category="error")
821824
log.error(e)
825+
return redirect(url_for("google.login"))
822826
return redirect(url_for('web.login'))
823827

824828

@@ -841,10 +845,6 @@ def generic_login():
841845
del oauthblueprints[2]['blueprint'].token
842846
except Exception:
843847
pass
844-
# Force clear session to prevent loops
845-
if 'generic_oauth_token' in session: session.pop('generic_oauth_token', None)
846-
if 'generic_oauth_user_id' in session: session.pop('generic_oauth_user_id', None)
847-
848848
flash(_("OAuth error: {}").format(e), category="error")
849849
log.error(e)
850850
return redirect(url_for("generic.login"))
@@ -853,12 +853,9 @@ def generic_login():
853853
del oauthblueprints[2]['blueprint'].token
854854
except Exception:
855855
pass
856-
# Force clear session to prevent loops
857-
if 'generic_oauth_token' in session: session.pop('generic_oauth_token', None)
858-
if 'generic_oauth_user_id' in session: session.pop('generic_oauth_user_id', None)
859-
860856
flash(_("OAuth error: {}").format(e), category="error")
861857
log.error(e)
858+
return redirect(url_for("generic.login"))
862859
return redirect(url_for("web.login"))
863860

864861

0 commit comments

Comments
 (0)