sadpandajoe commented on code in PR #38469:
URL: https://github.com/apache/superset/pull/38469#discussion_r3828166357


##########
superset/utils/oauth2.py:
##########
@@ -318,3 +318,186 @@ def check_for_oauth2(database: Database) -> 
Iterator[None]:
         if database.is_oauth2_enabled() and 
database.db_engine_spec.needs_oauth2(ex):
             database.db_engine_spec.start_oauth2_dance(database)
         raise
+
+
+def get_access_token_for_database(database: Database, user_id: int) -> str | 
None:
+    """
+    Return a valid OAuth2 access token for the given database and user.
+
+    First checks if the database has an upstream OAuth provider configured in 
its
+    ``extra`` JSON (key ``oauth2_upstream_provider``). If so, returns the saved
+    upstream login token for that provider.
+
+    Otherwise, falls back to the database-specific OAuth2 flow.
+    """
+    upstream_provider = database.get_extra().get("oauth2_upstream_provider")
+    if upstream_provider:
+        access_token = get_upstream_provider_token(upstream_provider, user_id)
+        if access_token:
+            logger.info(
+                "Using upstream OAuth token from provider '%s' "
+                "for database '%s' (id=%s, user_id=%d)",
+                upstream_provider,
+                database.database_name,
+                database.id,
+                user_id,
+            )
+            return access_token
+        logger.warning(
+            "Upstream provider '%s' configured for database '%s' "
+            "(id=%s) but no valid token found for user_id=%d, "
+            "falling back to database-specific OAuth2",
+            upstream_provider,
+            database.database_name,
+            database.id,
+            user_id,
+        )
+
+    # Fall back to database-specific OAuth2 (also used when upstream token is
+    # unavailable, e.g. expired without a refresh token).
+    # Database-specific OAuth2 requires a persisted database (needs database.id
+    # to look up stored tokens).
+    if database.id is None:
+        logger.debug(
+            "Database '%s' has no persisted id, "
+            "skipping database-specific OAuth2 token lookup",
+            database.database_name,
+        )
+        return None
+
+    oauth2_config = database.get_oauth2_config()
+    if oauth2_config:
+        logger.info(
+            "Using database-specific OAuth2 token "
+            "for database '%s' (id=%d, user_id=%d)",
+            database.database_name,
+            database.id,
+            user_id,
+        )
+        return get_oauth2_access_token(
+            oauth2_config, database.id, user_id, database.db_engine_spec
+        )
+
+    logger.warning(
+        "No OAuth2 token available for database '%s' (id=%d, user_id=%d)",
+        database.database_name,
+        database.id,
+        user_id,
+    )
+    return None
+
+
+def save_user_provider_token(
+    user_id: int,
+    provider: str,
+    token_response: dict[str, Any],
+) -> None:
+    """
+    Upsert an UpstreamOAuthToken row for the given user + provider.
+    """
+    from superset.models.core import UpstreamOAuthToken
+
+    token: UpstreamOAuthToken | None = (
+        db.session.query(UpstreamOAuthToken)
+        .filter_by(user_id=user_id, provider=provider)
+        .one_or_none()
+    )
+    if token is None:
+        token = UpstreamOAuthToken(user_id=user_id, provider=provider)
+
+    token.access_token = token_response.get("access_token")

Review Comment:
   FAB supports a provider-specific `token_key`, but this persistence path 
always reads `access_token`. Providers configured with another token field can 
log in successfully while this stores `None`, so upstream forwarding never 
receives their token. Could this resolve the configured token key before 
persisting it?



##########
superset/utils/oauth2.py:
##########
@@ -318,3 +318,186 @@ def check_for_oauth2(database: Database) -> 
Iterator[None]:
         if database.is_oauth2_enabled() and 
database.db_engine_spec.needs_oauth2(ex):
             database.db_engine_spec.start_oauth2_dance(database)
         raise
+
+
+def get_access_token_for_database(database: Database, user_id: int) -> str | 
None:
+    """
+    Return a valid OAuth2 access token for the given database and user.
+
+    First checks if the database has an upstream OAuth provider configured in 
its
+    ``extra`` JSON (key ``oauth2_upstream_provider``). If so, returns the saved
+    upstream login token for that provider.
+
+    Otherwise, falls back to the database-specific OAuth2 flow.
+    """
+    upstream_provider = database.get_extra().get("oauth2_upstream_provider")
+    if upstream_provider:
+        access_token = get_upstream_provider_token(upstream_provider, user_id)
+        if access_token:
+            logger.info(
+                "Using upstream OAuth token from provider '%s' "
+                "for database '%s' (id=%s, user_id=%d)",
+                upstream_provider,
+                database.database_name,
+                database.id,
+                user_id,
+            )
+            return access_token
+        logger.warning(
+            "Upstream provider '%s' configured for database '%s' "
+            "(id=%s) but no valid token found for user_id=%d, "
+            "falling back to database-specific OAuth2",
+            upstream_provider,
+            database.database_name,
+            database.id,
+            user_id,
+        )
+
+    # Fall back to database-specific OAuth2 (also used when upstream token is
+    # unavailable, e.g. expired without a refresh token).
+    # Database-specific OAuth2 requires a persisted database (needs database.id
+    # to look up stored tokens).
+    if database.id is None:
+        logger.debug(
+            "Database '%s' has no persisted id, "
+            "skipping database-specific OAuth2 token lookup",
+            database.database_name,
+        )
+        return None
+
+    oauth2_config = database.get_oauth2_config()
+    if oauth2_config:
+        logger.info(
+            "Using database-specific OAuth2 token "
+            "for database '%s' (id=%d, user_id=%d)",
+            database.database_name,
+            database.id,
+            user_id,
+        )
+        return get_oauth2_access_token(
+            oauth2_config, database.id, user_id, database.db_engine_spec
+        )
+
+    logger.warning(

Review Comment:
   `_get_sqla_engine()` now calls this helper for every authenticated engine 
creation, including databases that do not use OAuth. That normal path reaches 
this warning on every query/chart render, flooding production logs with a false 
failure signal. Should the no-OAuth case return quietly (and reserve warnings 
for configured OAuth failures)?



##########
superset/utils/oauth2.py:
##########
@@ -318,3 +318,186 @@ def check_for_oauth2(database: Database) -> 
Iterator[None]:
         if database.is_oauth2_enabled() and 
database.db_engine_spec.needs_oauth2(ex):
             database.db_engine_spec.start_oauth2_dance(database)
         raise
+
+
+def get_access_token_for_database(database: Database, user_id: int) -> str | 
None:
+    """
+    Return a valid OAuth2 access token for the given database and user.
+
+    First checks if the database has an upstream OAuth provider configured in 
its
+    ``extra`` JSON (key ``oauth2_upstream_provider``). If so, returns the saved
+    upstream login token for that provider.
+
+    Otherwise, falls back to the database-specific OAuth2 flow.
+    """
+    upstream_provider = database.get_extra().get("oauth2_upstream_provider")
+    if upstream_provider:
+        access_token = get_upstream_provider_token(upstream_provider, user_id)
+        if access_token:
+            logger.info(
+                "Using upstream OAuth token from provider '%s' "
+                "for database '%s' (id=%s, user_id=%d)",
+                upstream_provider,
+                database.database_name,
+                database.id,
+                user_id,
+            )
+            return access_token
+        logger.warning(
+            "Upstream provider '%s' configured for database '%s' "
+            "(id=%s) but no valid token found for user_id=%d, "
+            "falling back to database-specific OAuth2",
+            upstream_provider,
+            database.database_name,
+            database.id,
+            user_id,
+        )
+
+    # Fall back to database-specific OAuth2 (also used when upstream token is
+    # unavailable, e.g. expired without a refresh token).
+    # Database-specific OAuth2 requires a persisted database (needs database.id
+    # to look up stored tokens).
+    if database.id is None:
+        logger.debug(
+            "Database '%s' has no persisted id, "
+            "skipping database-specific OAuth2 token lookup",
+            database.database_name,
+        )
+        return None
+
+    oauth2_config = database.get_oauth2_config()
+    if oauth2_config:
+        logger.info(
+            "Using database-specific OAuth2 token "
+            "for database '%s' (id=%d, user_id=%d)",
+            database.database_name,
+            database.id,
+            user_id,
+        )
+        return get_oauth2_access_token(
+            oauth2_config, database.id, user_id, database.db_engine_spec
+        )
+
+    logger.warning(
+        "No OAuth2 token available for database '%s' (id=%d, user_id=%d)",
+        database.database_name,
+        database.id,
+        user_id,
+    )
+    return None
+
+
+def save_user_provider_token(
+    user_id: int,
+    provider: str,
+    token_response: dict[str, Any],
+) -> None:
+    """
+    Upsert an UpstreamOAuthToken row for the given user + provider.
+    """
+    from superset.models.core import UpstreamOAuthToken
+
+    token: UpstreamOAuthToken | None = (
+        db.session.query(UpstreamOAuthToken)
+        .filter_by(user_id=user_id, provider=provider)
+        .one_or_none()
+    )
+    if token is None:
+        token = UpstreamOAuthToken(user_id=user_id, provider=provider)
+
+    token.access_token = token_response.get("access_token")
+    expires_in = token_response.get("expires_in")
+    token.access_token_expiration = (
+        datetime.now() + timedelta(seconds=expires_in) if expires_in else None
+    )
+    token.refresh_token = token_response.get("refresh_token")
+    db.session.add(token)
+    db.session.commit()
+
+
[email protected]_exception(
+    backoff.expo,
+    AcquireDistributedLockFailedException,
+    factor=10,
+    base=2,
+    max_tries=5,
+)
+def get_upstream_provider_token(provider: str, user_id: int) -> str | None:
+    """
+    Retrieve a valid access token for the given provider and user.
+
+    If the token is expired and a refresh token exists, attempt to refresh it.
+    Returns None if no valid token is available.
+    """
+    from superset.models.core import UpstreamOAuthToken
+
+    token: UpstreamOAuthToken | None = (
+        db.session.query(UpstreamOAuthToken)
+        .filter_by(user_id=user_id, provider=provider)
+        .one_or_none()
+    )
+    if token is None:
+        return None
+
+    now = datetime.now()
+    if token.access_token_expiration is None or token.access_token_expiration 
> now:
+        return token.access_token
+
+    # Token is expired
+    if token.refresh_token:
+        return _refresh_upstream_provider_token(token, provider)
+
+    db.session.delete(token)
+    db.session.commit()
+    return None
+
+
+def _refresh_upstream_provider_token(
+    token: UpstreamOAuthToken,
+    provider: str,
+) -> str | None:
+    """
+    Use the refresh token to obtain a new access token from the provider.
+    Updates and persists the token if successful; deletes it on failure.
+    """
+    from flask import current_app as flask_app
+
+    with DistributedLock(
+        namespace="refresh_upstream_oauth_token",
+        user_id=token.user_id,
+        provider=provider,
+    ):
+        try:
+            remote_app = 
flask_app.extensions["authlib.integrations.flask_client"][

Review Comment:
   This extension entry is Authlib's OAuth registry, not a mapping of 
configured remotes, so subscripting it raises before `fetch_access_token()` is 
reached. The broad handler then deletes the stored token; every expired 
upstream token with a refresh token is therefore lost. Should this use the 
configured OAuth remote accessor instead?



##########
superset/utils/oauth2.py:
##########
@@ -318,3 +318,186 @@ def check_for_oauth2(database: Database) -> 
Iterator[None]:
         if database.is_oauth2_enabled() and 
database.db_engine_spec.needs_oauth2(ex):
             database.db_engine_spec.start_oauth2_dance(database)
         raise
+
+
+def get_access_token_for_database(database: Database, user_id: int) -> str | 
None:
+    """
+    Return a valid OAuth2 access token for the given database and user.
+
+    First checks if the database has an upstream OAuth provider configured in 
its
+    ``extra`` JSON (key ``oauth2_upstream_provider``). If so, returns the saved
+    upstream login token for that provider.
+
+    Otherwise, falls back to the database-specific OAuth2 flow.
+    """
+    upstream_provider = database.get_extra().get("oauth2_upstream_provider")
+    if upstream_provider:
+        access_token = get_upstream_provider_token(upstream_provider, user_id)
+        if access_token:
+            logger.info(
+                "Using upstream OAuth token from provider '%s' "
+                "for database '%s' (id=%s, user_id=%d)",
+                upstream_provider,
+                database.database_name,
+                database.id,
+                user_id,
+            )
+            return access_token
+        logger.warning(
+            "Upstream provider '%s' configured for database '%s' "
+            "(id=%s) but no valid token found for user_id=%d, "
+            "falling back to database-specific OAuth2",
+            upstream_provider,
+            database.database_name,
+            database.id,
+            user_id,
+        )
+
+    # Fall back to database-specific OAuth2 (also used when upstream token is
+    # unavailable, e.g. expired without a refresh token).
+    # Database-specific OAuth2 requires a persisted database (needs database.id
+    # to look up stored tokens).
+    if database.id is None:
+        logger.debug(
+            "Database '%s' has no persisted id, "
+            "skipping database-specific OAuth2 token lookup",
+            database.database_name,
+        )
+        return None
+
+    oauth2_config = database.get_oauth2_config()
+    if oauth2_config:
+        logger.info(
+            "Using database-specific OAuth2 token "
+            "for database '%s' (id=%d, user_id=%d)",
+            database.database_name,
+            database.id,
+            user_id,
+        )
+        return get_oauth2_access_token(
+            oauth2_config, database.id, user_id, database.db_engine_spec
+        )
+
+    logger.warning(
+        "No OAuth2 token available for database '%s' (id=%d, user_id=%d)",
+        database.database_name,
+        database.id,
+        user_id,
+    )
+    return None
+
+
+def save_user_provider_token(
+    user_id: int,
+    provider: str,
+    token_response: dict[str, Any],
+) -> None:
+    """
+    Upsert an UpstreamOAuthToken row for the given user + provider.
+    """
+    from superset.models.core import UpstreamOAuthToken
+
+    token: UpstreamOAuthToken | None = (
+        db.session.query(UpstreamOAuthToken)
+        .filter_by(user_id=user_id, provider=provider)
+        .one_or_none()
+    )
+    if token is None:
+        token = UpstreamOAuthToken(user_id=user_id, provider=provider)
+
+    token.access_token = token_response.get("access_token")
+    expires_in = token_response.get("expires_in")
+    token.access_token_expiration = (
+        datetime.now() + timedelta(seconds=expires_in) if expires_in else None
+    )
+    token.refresh_token = token_response.get("refresh_token")

Review Comment:
   A second authorization-code login is allowed to return a new access token 
without a `refresh_token`. This assignment then erases the still-valid stored 
refresh token, so the next expiry cannot be refreshed and the forwarding flow 
falls back to another OAuth dance. Should this preserve the existing refresh 
token unless the response explicitly supplies a replacement?



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to