This commit is contained in:
Timothy Jaeryang Baek
2026-04-13 15:13:14 -05:00
parent 9c64d84ad9
commit 31406caa79
2 changed files with 43 additions and 32 deletions
+2 -2
View File
@@ -123,7 +123,7 @@ class OAuthSessionTable:
'user_id': user_id,
'provider': provider,
'token': self._encrypt_token(token),
'expires_at': token.get('expires_at'),
'expires_at': token.get('expires_at') or int(time.time() + 3600),
'created_at': current_time,
'updated_at': current_time,
}
@@ -274,7 +274,7 @@ class OAuthSessionTable:
.filter_by(id=session_id)
.values(
token=self._encrypt_token(token),
expires_at=token.get('expires_at'),
expires_at=token.get('expires_at') or int(time.time() + 3600),
updated_at=current_time,
)
)
+41 -30
View File
@@ -140,6 +140,41 @@ auth_manager_config.OAUTH_UPDATE_EMAIL_ON_LOGIN = OAUTH_UPDATE_EMAIL_ON_LOGIN
auth_manager_config.OAUTH_AUDIENCE = OAUTH_AUDIENCE
# Conservative default when the provider omits both expires_in and expires_at.
# Matches the value recommended by Authlib's compliance_fix documentation.
DEFAULT_TOKEN_EXPIRY_SECONDS = 3600
def _normalize_token_expiry(token: dict) -> dict:
"""Ensure a token dict always has a numeric ``expires_at``.
Resolution order:
1. If *expires_at* is already present and non-None, trust it.
2. Else if *expires_in* is present and non-None, compute *expires_at*.
3. Otherwise fall back to ``DEFAULT_TOKEN_EXPIRY_SECONDS`` and log a
warning so operators can identify providers that omit expiration.
Also stamps *issued_at* for auditing.
"""
token['issued_at'] = datetime.now().timestamp()
if token.get('expires_at') is not None:
token['expires_at'] = int(token['expires_at'])
return token
if token.get('expires_in') is not None:
token['expires_at'] = int(datetime.now().timestamp() + token['expires_in'])
return token
# Neither field present — conservative fallback
log.warning(
"OAuth token response missing both 'expires_in' and 'expires_at'; "
f"defaulting to {DEFAULT_TOKEN_EXPIRY_SECONDS}s from now"
)
token['expires_at'] = int(datetime.now().timestamp() + DEFAULT_TOKEN_EXPIRY_SECONDS)
return token
FERNET = None
if len(OAUTH_CLIENT_INFO_ENCRYPTION_KEY) != 44:
@@ -712,7 +747,7 @@ class OAuthClientManager:
log.warning(f'No OAuth session found for user {user_id}, client_id {client_id}')
return None
if force_refresh or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at):
if force_refresh or session.expires_at is None or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at):
log.debug(f'Token refresh needed for user {user_id}, client_id {session.provider}')
refreshed_token = await self._refresh_token(session)
if refreshed_token:
@@ -823,14 +858,7 @@ class OAuthClientManager:
if 'refresh_token' not in new_token_data:
new_token_data['refresh_token'] = token_data['refresh_token']
# Add timestamp for tracking
new_token_data['issued_at'] = datetime.now().timestamp()
# Calculate expires_at if we have expires_in
if 'expires_in' in new_token_data and 'expires_at' not in new_token_data:
new_token_data['expires_at'] = int(
datetime.now().timestamp() + new_token_data['expires_in']
)
_normalize_token_expiry(new_token_data)
log.debug(f'Token refresh successful for client_id {client_id}')
return new_token_data
@@ -883,12 +911,7 @@ class OAuthClientManager:
if token:
try:
# Add timestamp for tracking
token['issued_at'] = datetime.now().timestamp()
# Calculate expires_at if we have expires_in
if 'expires_in' in token and 'expires_at' not in token:
token['expires_at'] = datetime.now().timestamp() + token['expires_in']
_normalize_token_expiry(token)
# Clean up any existing sessions for this user/client_id first
sessions = await OAuthSessions.get_sessions_by_user_id(user_id)
@@ -975,7 +998,7 @@ class OAuthManager:
log.warning(f'No OAuth session found for user {user_id}, session {session_id}')
return None
if force_refresh or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at):
if force_refresh or session.expires_at is None or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at):
log.debug(f'Token refresh needed for user {user_id}, provider {session.provider}')
refreshed_token = await self._refresh_token(session)
if refreshed_token:
@@ -1089,14 +1112,7 @@ class OAuthManager:
if 'refresh_token' not in new_token_data:
new_token_data['refresh_token'] = token_data['refresh_token']
# Add timestamp for tracking
new_token_data['issued_at'] = datetime.now().timestamp()
# Calculate expires_at if we have expires_in
if 'expires_in' in new_token_data and 'expires_at' not in new_token_data:
new_token_data['expires_at'] = int(
datetime.now().timestamp() + new_token_data['expires_in']
)
_normalize_token_expiry(new_token_data)
log.debug(f'Token refresh successful for provider {provider}')
return new_token_data
@@ -1694,12 +1710,7 @@ class OAuthManager:
)
try:
# Add timestamp for tracking
token['issued_at'] = datetime.now().timestamp()
# Calculate expires_at if we have expires_in
if 'expires_in' in token and 'expires_at' not in token:
token['expires_at'] = datetime.now().timestamp() + token['expires_in']
_normalize_token_expiry(token)
# Enforce max concurrent sessions per user/provider to prevent
# unbounded growth while allowing multi-device usage