diff --git a/web/pgadmin/authenticate/oauth2.py b/web/pgadmin/authenticate/oauth2.py index 4e04b5b23be..e4328e7575f 100644 --- a/web/pgadmin/authenticate/oauth2.py +++ b/web/pgadmin/authenticate/oauth2.py @@ -687,6 +687,8 @@ def get_user_profile(self): self.oauth2_current_client, provider, client ) + session["oauth2_provider"] = self.oauth2_current_client + session['pass_enc_key'] = session['oauth2_token']['access_token'] if 'OAUTH2_LOGOUT_URL' in self.oauth2_config[ diff --git a/web/pgadmin/browser/server_groups/servers/__init__.py b/web/pgadmin/browser/server_groups/servers/__init__.py index cfcb324c4d8..c0488f94f70 100644 --- a/web/pgadmin/browser/server_groups/servers/__init__.py +++ b/web/pgadmin/browser/server_groups/servers/__init__.py @@ -1636,6 +1636,13 @@ def connect(self, gid, sid, is_qt=False, server=None): manager.passexec = None conn = manager.connection() + connection_params = manager.connection_params or {} + + use_pgadmin_oauth = connection_params.get( + "oauth_pgadmin_token_mode", + "disabled", + ) in ("direct", "exchange") + # Get enc key crypt_key_present, crypt_key = get_crypt_key() if not crypt_key_present: @@ -1664,8 +1671,13 @@ def connect(self, gid, sid, is_qt=False, server=None): except Exception as e: current_app.logger.exception(e) return internal_server_error(errormsg=str(e)) - if 'password' not in data and (server.kerberos_conn is False or - server.kerberos_conn is None): + + if use_pgadmin_oauth: + # The pgAdmin OAuth bearer token is the database credential. + password = None + save_password = False + elif 'password' not in data and (server.kerberos_conn is False or + server.kerberos_conn is None): passfile_param = None if hasattr(server, 'connection_params') and \ @@ -1716,9 +1728,19 @@ def connect(self, gid, sid, is_qt=False, server=None): server_types=ServerType.types() ) except Exception as e: + error_message = getattr(e, 'message', str(e)) + + if use_pgadmin_oauth: + return make_json_response( + status=400, + success=0, + errormsg=error_message + ) + return self.get_response_for_password( server, 401, not server.save_password, prompt_tunnel_password, - getattr(e, 'message', str(e))) + error_message + ) if not status: current_app.logger.error( @@ -1728,6 +1750,13 @@ def connect(self, gid, sid, is_qt=False, server=None): if errmsg.find('Ticket expired') != -1: return internal_server_error(errmsg) + if use_pgadmin_oauth: + return make_json_response( + status=400, + success=0, + errormsg=errmsg + ) + return self.get_response_for_password( server, 401, not server.save_password, prompt_tunnel_password, errmsg) diff --git a/web/pgadmin/browser/server_groups/servers/static/js/server.ui.js b/web/pgadmin/browser/server_groups/servers/static/js/server.ui.js index 47d62cfca9b..b7de8ad65e7 100644 --- a/web/pgadmin/browser/server_groups/servers/static/js/server.ui.js +++ b/web/pgadmin/browser/server_groups/servers/static/js/server.ui.js @@ -165,6 +165,12 @@ export function getConnectionParameters() { }, { 'value': 'oauth_scope', 'label': gettext('OAuth scope'), 'vartype': 'string', 'min_server_version': '18' + }, { + 'value': 'oauth_pgadmin_token_mode', + 'label': gettext('OAuth pgAdmin token mode'), + 'vartype': 'enum', + 'enumvals': ['disabled', 'direct', 'exchange'], + 'min_server_version': '18' }]; conParams.sort(function (a, b) { diff --git a/web/pgadmin/utils/driver/psycopg3/connection.py b/web/pgadmin/utils/driver/psycopg3/connection.py index d07a16cefcd..2c57ca16618 100644 --- a/web/pgadmin/utils/driver/psycopg3/connection.py +++ b/web/pgadmin/utils/driver/psycopg3/connection.py @@ -43,6 +43,12 @@ from io import StringIO from pgadmin.utils.locker import ConnectionLocker from pgadmin.utils.driver import get_driver +from pgadmin.utils.pg_oauth2 import ( + install_oauth_hook, + get_postgres_oauth_token, + oauth_token_context, + OAuthTokenError, +) # On Windows, Psycopg is not compatible with the default ProactorEventLoop. @@ -362,25 +368,51 @@ def connect(self, **kwargs): connection_string = manager.create_connection_string( database, user, password) - if self.async_: - autocommit = True - if 'auto_commit' in kwargs: - autocommit = kwargs['auto_commit'] + connection_params = manager.connection_params or {} - async def connectdbserver(): - return await psycopg.AsyncConnection.connect( - connection_string, - cursor_factory=AsyncDictCursor, - autocommit=autocommit, - prepare_threshold=manager.prepare_threshold + oauth_mode = connection_params.get( + "oauth_pgadmin_token_mode", + "disabled", + ) + + oauth_token = None + + if oauth_mode != "disabled": + try: + oauth_token = get_postgres_oauth_token( + oauth_mode, + manager.connection_params.get("oauth_client_id"), ) - pg_conn = asyncio.run(connectdbserver()) - pg_conn.server_cursor_factory = AsyncDictServerCursor - else: - pg_conn = psycopg.Connection.connect( - connection_string, - cursor_factory=DictCursor, - prepare_threshold=manager.prepare_threshold) + install_oauth_hook() + except OAuthTokenError as exc: + current_app.logger.warning( + "PostgreSQL OAuth authentication " + "failed for server %s: %s", + manager.sid, + exc, + ) + return False, str(exc) + + with oauth_token_context(oauth_token): + if self.async_: + autocommit = True + if 'auto_commit' in kwargs: + autocommit = kwargs['auto_commit'] + + async def connectdbserver(): + return await psycopg.AsyncConnection.connect( + connection_string, + cursor_factory=AsyncDictCursor, + autocommit=autocommit, + prepare_threshold=manager.prepare_threshold + ) + pg_conn = asyncio.run(connectdbserver()) + pg_conn.server_cursor_factory = AsyncDictServerCursor + else: + pg_conn = psycopg.Connection.connect( + connection_string, + cursor_factory=DictCursor, + prepare_threshold=manager.prepare_threshold) except psycopg.Error as e: manager.stop_ssh_tunnel() diff --git a/web/pgadmin/utils/driver/psycopg3/server_manager.py b/web/pgadmin/utils/driver/psycopg3/server_manager.py index 00738355ee4..d0c0edb3bbe 100644 --- a/web/pgadmin/utils/driver/psycopg3/server_manager.py +++ b/web/pgadmin/utils/driver/psycopg3/server_manager.py @@ -677,6 +677,10 @@ def create_connection_string(self, database, user, password=None): # Loop through all the connection parameters set in the server dialog. if self.connection_params and isinstance(self.connection_params, dict): for key, value in self.connection_params.items(): + # pgAdmin-only parameter, not a libpq connection option + if key == "oauth_pgadmin_token_mode": + continue + with_complete_path = False orig_value = value # Getting complete file path if the key is one of the below. diff --git a/web/pgadmin/utils/driver/psycopg3/tests/test_connection_oauth.py b/web/pgadmin/utils/driver/psycopg3/tests/test_connection_oauth.py new file mode 100644 index 00000000000..02055762e0e --- /dev/null +++ b/web/pgadmin/utils/driver/psycopg3/tests/test_connection_oauth.py @@ -0,0 +1,519 @@ +########################################################################## +# +# pgAdmin 4 - PostgreSQL Tools +# +# Copyright (C) 2013 - 2026, The pgAdmin Development Team +# This software is released under the PostgreSQL Licence +# +########################################################################## + +"""Tests for pgAdmin OAuth integration in the Psycopg 3 connection.""" + +import unittest +from contextlib import ExitStack, nullcontext +from unittest.mock import AsyncMock, MagicMock, patch + +import psycopg +from flask import Flask + +from pgadmin.utils import pg_oauth2 +import pgadmin.utils.driver.psycopg3.connection as connection_module + + +class TestOAuthConnection(unittest.TestCase): + """Test pgAdmin OAuth handling during Psycopg connection creation.""" + + def setUp(self): + self.app = Flask(__name__) + self.app.config.update( + SECRET_KEY='connection-oauth-unit-test-secret', + TESTING=True, + ) + + self.token_handle = pg_oauth2._oauth_token.set(None) + + def tearDown(self): + pg_oauth2._oauth_token.reset(self.token_handle) + + @staticmethod + def _make_manager(mode='disabled', client_id=None): + manager = MagicMock() + + manager.sid = 1 + manager.db = 'postgres' + manager.user = 'database-user' + manager.password = None + manager.role = None + + manager.use_ssh_tunnel = 0 + manager.tunnel_created = False + manager.kerberos_conn = False + + manager.passexec = None + manager.prepare_threshold = None + + manager.connection_params = { + 'oauth_pgadmin_token_mode': mode, + } + + if client_id is not None: + manager.connection_params['oauth_client_id'] = client_id + + manager.get_connection_param_value.return_value = None + manager.create_connection_string.return_value = ( + 'host=database.example.test ' + 'dbname=postgres ' + 'user=database-user' + ) + + return manager + + @staticmethod + def _make_connection(manager, async_=False): + # Construct the object without invoking Connection.__init__(). + # Its constructor obtains unrelated global driver state that is + # not part of these tests. + connection = object.__new__(connection_module.Connection) + + connection.manager = manager + connection.conn_id = 'DB:postgres' + connection.db = 'postgres' + connection.conn = None + + connection.password = None + connection.reconnecting = False + connection.wasConnected = False + connection.async_ = 1 if async_ else 0 + + # Connection initialization executes several SQL statements. + # These tests stop at the Psycopg connection boundary. + connection._initialize = MagicMock( + return_value=(True, None) + ) + + return connection + + @staticmethod + def _make_psycopg_connection(): + pg_connection = MagicMock() + pg_connection.closed = False + return pg_connection + + @staticmethod + def _apply_common_patches(stack): + stack.enter_context( + patch.object( + connection_module, + 'get_crypt_key', + return_value=(True, 'crypt-key') + ) + ) + stack.enter_context( + patch.object( + connection_module, + 'get_complete_file_path', + return_value=None + ) + ) + stack.enter_context( + patch.object( + connection_module, + 'ConnectionLocker', + return_value=nullcontext() + ) + ) + stack.enter_context( + patch.object( + connection_module, + 'gettext', + side_effect=lambda message, *args, **kwargs: message + ) + ) + stack.enter_context( + patch.object( + connection_module, + '_', + side_effect=lambda message, *args, **kwargs: message + ) + ) + + def test_oauth_disabled_does_not_access_pgadmin_token(self): + """ + Existing non-OAuth connections must not install the OAuth hook + or access the pgAdmin OAuth session. + """ + manager = self._make_manager(mode='disabled') + connection = self._make_connection(manager) + pg_connection = self._make_psycopg_connection() + + def connect_without_oauth(*args, **kwargs): + self.assertIsNone(pg_oauth2._oauth_token.get()) + return pg_connection + + with self.app.test_request_context('/'): + with ExitStack() as stack: + self._apply_common_patches(stack) + + install_hook = stack.enter_context( + patch.object( + connection_module, + 'install_oauth_hook' + ) + ) + get_token = stack.enter_context( + patch.object( + connection_module, + 'get_postgres_oauth_token' + ) + ) + psycopg_connect = stack.enter_context( + patch.object( + connection_module.psycopg.Connection, + 'connect', + side_effect=connect_without_oauth + ) + ) + + status, message = connection.connect() + + self.assertTrue(status) + self.assertIsNone(message) + + install_hook.assert_not_called() + get_token.assert_not_called() + psycopg_connect.assert_called_once() + + connection._initialize.assert_called_once() + manager.create_connection_string.assert_called_once_with( + 'postgres', + 'database-user', + None + ) + + def test_oauth_enabled_without_token_returns_error(self): + """ + An OAuth-enabled server must fail clearly before calling + Psycopg when the pgAdmin session has no access token. + """ + manager = self._make_manager(mode='direct') + connection = self._make_connection(manager) + + with self.app.test_request_context('/'): + with ExitStack() as stack: + self._apply_common_patches(stack) + + install_hook = stack.enter_context( + patch.object( + connection_module, + 'install_oauth_hook' + ) + ) + get_token = stack.enter_context( + patch.object( + connection_module, + 'get_postgres_oauth_token', + side_effect=pg_oauth2.OAuthTokenError( + 'No current pgAdmin OAuth access token' + ' is available. Sign in to pgAdmin again.' + ) + ) + ) + psycopg_connect = stack.enter_context( + patch.object( + connection_module.psycopg.Connection, + 'connect' + ) + ) + + status, message = connection.connect() + + self.assertFalse(status) + self.assertEqual( + message, + 'No current pgAdmin OAuth access token ' + 'is available. Sign in to pgAdmin again.' + ) + + install_hook.assert_not_called() + get_token.assert_called_once_with('direct', None) + psycopg_connect.assert_not_called() + connection._initialize.assert_not_called() + + def test_sync_connection_is_created_inside_token_context(self): + """ + The synchronous Psycopg connection must be created while the + pgAdmin access token is present in the ContextVar. + """ + manager = self._make_manager(mode='direct') + connection = self._make_connection(manager) + pg_connection = self._make_psycopg_connection() + + def connect_with_oauth(*args, **kwargs): + self.assertEqual( + pg_oauth2._oauth_token.get(), + 'pgadmin-access-token' + ) + return pg_connection + + with self.app.test_request_context('/'): + with ExitStack() as stack: + self._apply_common_patches(stack) + + install_hook = stack.enter_context( + patch.object( + connection_module, + 'install_oauth_hook' + ) + ) + get_token = stack.enter_context( + patch.object( + connection_module, + 'get_postgres_oauth_token', + return_value='pgadmin-access-token' + ) + ) + psycopg_connect = stack.enter_context( + patch.object( + connection_module.psycopg.Connection, + 'connect', + side_effect=connect_with_oauth + ) + ) + + status, message = connection.connect() + + self.assertTrue(status) + self.assertIsNone(message) + self.assertIsNone(pg_oauth2._oauth_token.get()) + + install_hook.assert_called_once_with() + get_token.assert_called_once_with('direct', None) + psycopg_connect.assert_called_once() + + connection._initialize.assert_called_once() + self.assertIs(connection.conn, pg_connection) + + def test_async_connection_is_created_inside_token_context(self): + """ + The asynchronous Psycopg connection must inherit the token + context while asyncio.run() executes the connection coroutine. + """ + manager = self._make_manager(mode='direct') + connection = self._make_connection( + manager, + async_=True + ) + pg_connection = self._make_psycopg_connection() + + async def connect_with_oauth(*args, **kwargs): + self.assertEqual( + pg_oauth2._oauth_token.get(), + 'pgadmin-access-token' + ) + return pg_connection + + async_connect = AsyncMock( + side_effect=connect_with_oauth + ) + + with self.app.test_request_context('/'): + with ExitStack() as stack: + self._apply_common_patches(stack) + + install_hook = stack.enter_context( + patch.object( + connection_module, + 'install_oauth_hook' + ) + ) + get_token = stack.enter_context( + patch.object( + connection_module, + 'get_postgres_oauth_token', + return_value='pgadmin-access-token' + ) + ) + stack.enter_context( + patch.object( + connection_module.psycopg.AsyncConnection, + 'connect', + new=async_connect + ) + ) + + status, message = connection.connect() + + self.assertTrue(status) + self.assertIsNone(message) + self.assertIsNone(pg_oauth2._oauth_token.get()) + + install_hook.assert_called_once_with() + get_token.assert_called_once_with('direct', None) + async_connect.assert_awaited_once() + + connection._initialize.assert_called_once() + self.assertIs(connection.conn, pg_connection) + self.assertIs( + pg_connection.server_cursor_factory, + connection_module.AsyncDictServerCursor + ) + + def test_token_context_is_restored_after_connection_failure(self): + """ + A Psycopg failure must not leave the failed connection's token + available to later libpq callbacks. + """ + manager = self._make_manager(mode='direct') + connection = self._make_connection(manager) + + def fail_inside_oauth_context(*args, **kwargs): + self.assertEqual( + pg_oauth2._oauth_token.get(), + 'pgadmin-access-token' + ) + raise psycopg.OperationalError( + 'test connection failure' + ) + + with self.app.test_request_context('/'): + with ExitStack() as stack: + self._apply_common_patches(stack) + + stack.enter_context( + patch.object( + connection_module, + 'install_oauth_hook' + ) + ) + get_token = stack.enter_context( + patch.object( + connection_module, + 'get_postgres_oauth_token', + return_value='pgadmin-access-token' + ) + ) + psycopg_connect = stack.enter_context( + patch.object( + connection_module.psycopg.Connection, + 'connect', + side_effect=fail_inside_oauth_context + ) + ) + + status, message = connection.connect() + + self.assertFalse(status) + self.assertIn('test connection failure', message) + self.assertIsNone(pg_oauth2._oauth_token.get()) + + get_token.assert_called_once_with('direct', None) + psycopg_connect.assert_called_once() + connection._initialize.assert_not_called() + self.assertIsNone(connection.conn) + + def test_exchange_mode_passes_client_id(self): + """ + Token exchange mode must pass the configured OAuth client ID + when obtaining the PostgreSQL token. + """ + manager = self._make_manager( + mode='exchange', client_id='postgres-cluster') + connection = self._make_connection(manager) + pg_connection = self._make_psycopg_connection() + + def connect_with_exchanged_token(*args, **kwargs): + self.assertEqual( + pg_oauth2._oauth_token.get(), + 'exchanged-token' + ) + return pg_connection + + with self.app.test_request_context('/'): + with ExitStack() as stack: + self._apply_common_patches(stack) + + install_hook = stack.enter_context( + patch.object( + connection_module, + 'install_oauth_hook' + ) + ) + get_token = stack.enter_context( + patch.object( + connection_module, + 'get_postgres_oauth_token', + return_value='exchanged-token' + ) + ) + psycopg_connect = stack.enter_context( + patch.object( + connection_module.psycopg.Connection, + 'connect', + side_effect=connect_with_exchanged_token + ) + ) + + status, message = connection.connect() + + self.assertTrue(status) + self.assertIsNone(message) + + install_hook.assert_called_once_with() + get_token.assert_called_once_with('exchange', 'postgres-cluster') + psycopg_connect.assert_called_once() + connection._initialize.assert_called_once() + + def test_missing_connection_params_uses_non_oauth_path(self): + """ + A connection manager without connection parameters must behave as a + normal non-OAuth connection. + """ + manager = self._make_manager() + manager.connection_params = None + + connection = self._make_connection(manager) + pg_connection = self._make_psycopg_connection() + + def connect_without_oauth(*args, **kwargs): + self.assertIsNone(pg_oauth2._oauth_token.get()) + return pg_connection + + with self.app.test_request_context('/'): + with ExitStack() as stack: + self._apply_common_patches(stack) + + install_hook = stack.enter_context( + patch.object( + connection_module, + 'install_oauth_hook' + ) + ) + get_token = stack.enter_context( + patch.object( + connection_module, + 'get_postgres_oauth_token' + ) + ) + psycopg_connect = stack.enter_context( + patch.object( + connection_module.psycopg.Connection, + 'connect', + side_effect=connect_without_oauth + ) + ) + + status, message = connection.connect() + + self.assertTrue(status) + self.assertIsNone(message) + + install_hook.assert_not_called() + get_token.assert_not_called() + psycopg_connect.assert_called_once() + + self.assertIs(connection.conn, pg_connection) + connection._initialize.assert_called_once() + + +if __name__ == '__main__': + unittest.main() diff --git a/web/pgadmin/utils/pg_oauth2.py b/web/pgadmin/utils/pg_oauth2.py new file mode 100644 index 00000000000..6ca77ea62e9 --- /dev/null +++ b/web/pgadmin/utils/pg_oauth2.py @@ -0,0 +1,764 @@ +########################################################################## +# +# pgAdmin 4 - PostgreSQL Tools +# +# Copyright (C) 2013 - 2026, The pgAdmin Development Team +# This software is released under the PostgreSQL Licence +# +########################################################################## + +""" +PostgreSQL 18 libpq OAuth bearer-token integration for pgAdmin. + +The libpq OAuth hook is process-global. The OAuth access token is therefore +not obtained from Flask's request/session context inside the C callback. + +Instead, connection.py establishes the token in a ContextVar immediately +before creating a Psycopg connection. The libpq callback reads that +ContextVar when authentication requests a bearer token. +""" + +import base64 +import ctypes +import ctypes.util +import json +import logging +import re +import threading +import time +from collections.abc import Mapping +from contextlib import contextmanager +from contextvars import ContextVar +from typing import Any, Iterator, Optional +from urllib.parse import urlsplit + +import requests +from flask import has_request_context, session +from requests.auth import AuthBase + +_install_lock = threading.Lock() +logger = logging.getLogger(__name__) + + +# PostgreSQL 18: +# +# typedef enum +# { +# PQAUTHDATA_PROMPT_OAUTH_DEVICE, +# PQAUTHDATA_OAUTH_BEARER_TOKEN +# } PGauthData; +PQAUTHDATA_OAUTH_BEARER_TOKEN = 1 + +_OAUTH_REFRESH_LEEWAY_SECONDS = 30 + + +class PGoauthBearerRequest(ctypes.Structure): + """ + PostgreSQL 18 PGoauthBearerRequest. + + The async and cleanup members are function pointers in C. They are + represented as c_void_p here because we assign the corresponding + ctypes callbacks explicitly below. + """ + + _fields_ = [ + ("openid_configuration", ctypes.c_char_p), + ("scope", ctypes.c_char_p), + ("async_", ctypes.c_void_p), + ("cleanup", ctypes.c_void_p), + ("token", ctypes.c_char_p), + ("user", ctypes.c_void_p), + ] + + +_AUTH_HOOK = ctypes.CFUNCTYPE( + ctypes.c_int, + ctypes.c_int, # PGauthData + ctypes.c_void_p, # PGconn * + ctypes.c_void_p, # void * +) + + +_CLEANUP_HOOK = ctypes.CFUNCTYPE( + None, + ctypes.c_void_p, # PGconn * + ctypes.c_void_p, # PGoauthBearerRequest * +) + + +# The token is established by connection.py while a connection is being +# created. ContextVar is preferable to Flask's session here because the +# libpq callback is process-global and may be invoked from C code. +_oauth_token: ContextVar[Optional[str]] = ContextVar( + "pgadmin_oauth_token", + default=None, +) + + +# Keep native objects alive for as long as libpq can call them. +_libpq = None +_hook = None +_cleanup_hook = None +_previous_hook = None + +# PGoauthBearerRequest pointer -> token buffer. +# +# libpq owns the request structure and calls cleanup() when it no longer +# needs request->token. +_token_buffers = {} + + +def _delegate_to_previous(authdata_type, conn, data): + """Delegate auth-data requests not handled by pgAdmin.""" + if _previous_hook is None: + return 0 + + try: + return _previous_hook(authdata_type, conn, data) + except Exception: + # Never allow an exception to escape through a C callback. + logger.exception("Previous libpq auth-data hook failed.") + return -1 + + +def _get_libpq() -> ctypes.CDLL: + """ + Load the system libpq and configure the PostgreSQL OAuth hook API. + + pgAdmin's source installation uses psycopg[c], which links against + the system libpq. PostgreSQL 18 or newer is required. + """ + global _libpq + + if _libpq is not None: + return _libpq + + path = ctypes.util.find_library("pq") + + if not path: + raise RuntimeError( + "Could not find libpq. PostgreSQL 18 or newer is required " + "for pgAdmin PostgreSQL OAuth bearer-token support." + ) + + libpq = ctypes.CDLL(path) + + try: + pq_lib_version = libpq.PQlibVersion + except AttributeError: + raise RuntimeError("Loaded libpq does not expose PQlibVersion().") + + pq_lib_version.argtypes = [] + pq_lib_version.restype = ctypes.c_int + + version = pq_lib_version() + + if version < 180000: + raise RuntimeError( + "PostgreSQL 18 or newer libpq is required for OAuth " + "bearer-token authentication; found libpq version {0}.".format( + version) + ) + + try: + pq_set_auth_data_hook = libpq.PQsetAuthDataHook + pq_get_auth_data_hook = libpq.PQgetAuthDataHook + except AttributeError: + raise RuntimeError( + "The loaded libpq does not provide the PostgreSQL 18 " + "OAuth auth-data hook API." + ) + + pq_set_auth_data_hook.argtypes = [_AUTH_HOOK] + pq_set_auth_data_hook.restype = None + + pq_get_auth_data_hook.argtypes = [] + pq_get_auth_data_hook.restype = ctypes.c_void_p + + _libpq = libpq + + logger.info( + "Using libpq version %d for" + "PostgreSQL OAuth support.", version + ) + + return _libpq + + +@_CLEANUP_HOOK +def _oauth_cleanup(conn, request): + """ + Release native memory associated with a libpq OAuth request. + + libpq calls this after it no longer needs request->token. + """ + try: + request_key = int(ctypes.cast(request, ctypes.c_void_p).value) + + _token_buffers.pop(request_key, None) + + logger.debug( + "Released PostgreSQL OAuth bearer-token" + "buffer for request %s.", request_key + ) + except Exception: + # Never allow an exception to escape a C callback. + pass + + +@_AUTH_HOOK +def _oauth_hook(authdata_type, conn, data): + """ + Supply the pgAdmin OAuth access token when libpq requests one. + """ + if authdata_type != PQAUTHDATA_OAUTH_BEARER_TOKEN: + return _delegate_to_previous(authdata_type, conn, data) + + try: + token = _oauth_token.get() + + if not token: + # This may be an unrelated OAuth connection made after the + # process-global pgAdmin hook was installed. + return _delegate_to_previous(authdata_type, conn, data) + + request = ctypes.cast(data, ctypes.POINTER( + PGoauthBearerRequest)).contents + + token_buffer = ctypes.create_string_buffer(token.encode("utf-8")) + + request_key = int(ctypes.cast(data, ctypes.c_void_p).value) + + # Keep the token memory alive until libpq invokes cleanup(). + _token_buffers[request_key] = token_buffer + + request.token = ctypes.cast(token_buffer, ctypes.c_char_p) + + # The token is supplied synchronously. No async callback is needed. + request.async_ = None + + # Tell libpq when it is safe to release the token buffer. + request.cleanup = ctypes.cast(_oauth_cleanup, ctypes.c_void_p) + + logger.debug( + "Supplying OAuth bearer token to " + "PostgreSQL libpq (token length=%d).", + len(token), + ) + + return 1 + + except Exception: + # Never allow an exception to cross the C callback boundary. + logger.exception( + "Exception while supplying OAuth bearer token to libpq.") + return -1 + + +@contextmanager +def oauth_token_context(token: str) -> Iterator[None]: + """ + Associate an OAuth access token with the current connection operation. + + The token is available to the libpq callback through ContextVar and is + automatically restored after the connection attempt. + """ + token_handle = _oauth_token.set(token) + + try: + yield + finally: + _oauth_token.reset(token_handle) + + +def install_oauth_hook() -> None: + """ + Install the pgAdmin OAuth bearer-token hook into libpq. + + This should be called once during application initialization, after + Psycopg has been imported/initialised. + """ + global _hook + global _cleanup_hook + global _previous_hook + + if _hook is not None: + return + + with _install_lock: + if _hook is not None: + return + + libpq = _get_libpq() + + previous_hook_ptr = libpq.PQgetAuthDataHook() + _previous_hook = _AUTH_HOOK( + previous_hook_ptr) if previous_hook_ptr else None + + _hook = _oauth_hook + _cleanup_hook = _oauth_cleanup + + libpq.PQsetAuthDataHook(_hook) + + logger.info("Installed PostgreSQL libpq OAuth bearer-token hook.") + + +def _get_token_expiry(token_data: Mapping[str, Any]) -> Optional[float]: + """ + Return the access-token expiration timestamp. + + Authlib normally stores expires_at. Fall back to the JWT exp claim + because some providers omit expires_at from their token response. + """ + expires_at = token_data.get("expires_at") + + if expires_at is not None: + try: + return float(expires_at) + except (TypeError, ValueError): + logger.warning("Invalid expires_at value in pgAdmin OAuth token.") + + access_token = token_data.get("access_token") + + if not isinstance(access_token, str) or not access_token: + return None + + try: + parts = access_token.split(".") + + if len(parts) != 3: + # Opaque access token: expiration cannot be decoded locally. + return None + + payload = parts[1] + payload += "=" * (-len(payload) % 4) + + claims = json.loads(base64.urlsafe_b64decode(payload.encode("ascii"))) + + if not isinstance(claims, dict): + return None + + expires_at = claims.get("exp") + + if expires_at is not None: + return float(expires_at) + + except (ValueError, TypeError, UnicodeError, json.JSONDecodeError): + logger.warning( + "Could not decode expiration from pgAdmin OAuth access token.") + + return None + + +class OAuthTokenError(RuntimeError): + """Unable to obtain a PostgreSQL OAuth bearer token.""" + + +class OAuthTokenExchangeError(OAuthTokenError): + """Token exchange failed; the connection must not proceed.""" + + +class _TokenExchangeNoAuth(AuthBase): + """Prevent implicit HTTP Basic authentication from a .netrc file.""" + + def __call__( + self, + request: requests.PreparedRequest, + ) -> requests.PreparedRequest: + return request + + +def _get_current_oauth_client() -> Any: + """Return the OAuth client used for the current pgAdmin login.""" + if not has_request_context(): + raise OAuthTokenError( + "OAuth authentication requires a pgAdmin login session." + ) + + provider_name = session.get("oauth2_provider") + + if not isinstance(provider_name, str) or not provider_name: + raise OAuthTokenError( + "The OAuth provider is not recorded in this pgAdmin " + "session. Sign in to pgAdmin again." + ) + + # Import lazily to avoid an import cycle during pgAdmin startup. + from pgadmin.authenticate import get_auth_sources + from pgadmin.utils.constants import OAUTH2 + + oauth_source = get_auth_sources(OAUTH2) + + if oauth_source is None: + raise OAuthTokenError( + "The pgAdmin OAuth authentication source is unavailable." + ) + + oauth_client = oauth_source.oauth2_clients.get(provider_name) + + if oauth_client is None: + raise OAuthTokenError( + "The OAuth provider recorded in this pgAdmin session " + "is not registered." + ) + + return oauth_client + + +def _refresh_pgadmin_oauth_token( + token_data: Mapping[str, Any], +) -> Optional[str]: + """ + Refresh the pgAdmin OAuth access token. + + Return the refreshed access token, or None when refreshing is not + possible. + """ + refresh_token = token_data.get("refresh_token") + + if not isinstance(refresh_token, str) or not refresh_token: + logger.warning( + "The pgAdmin OAuth access token has expired, but no " + "refresh token is available." + ) + return None + + try: + oauth_client = _get_current_oauth_client() + + refreshed_token = oauth_client.fetch_access_token( + grant_type="refresh_token", + refresh_token=refresh_token, + ) + + if not isinstance(refreshed_token, dict): + logger.warning( + "The OAuth provider returned an invalid token " + "response while refreshing the pgAdmin access token." + ) + return None + + access_token = refreshed_token.get("access_token") + + if not isinstance(access_token, str) or not access_token: + logger.warning( + "The OAuth provider did not return a valid access " + "token while refreshing the pgAdmin access token." + ) + return None + + # Refresh responses may omit unchanged fields such as id_token, + # scope, token_type, or refresh_token. + updated_token = dict(token_data) + updated_token.update(refreshed_token) + + # Some providers do not return another refresh token. + # Keep the existing token in that case. + if not updated_token.get("refresh_token"): + updated_token["refresh_token"] = refresh_token + + session["oauth2_token"] = updated_token + session.modified = True + + logger.info("Refreshed pgAdmin OAuth access token.") + + return access_token + + except OAuthTokenError as exc: + logger.warning( + "Cannot refresh the pgAdmin OAuth access token: %s", + exc, + ) + return None + except Exception: + logger.exception( + "Failed to refresh the pgAdmin OAuth access token." + ) + return None + + +def _get_pgadmin_oauth_token() -> Optional[str]: + """ + Return the current pgAdmin OAuth access token. + + Refresh the token when it is expired or close to expiration. + Return None when no usable token is available. + """ + try: + token_data = session.get("oauth2_token") + except RuntimeError: + # No Flask request context. + return None + + if not isinstance(token_data, dict): + return None + + access_token = token_data.get("access_token") + + if not isinstance(access_token, str) or not access_token: + return None + + expires_at = _get_token_expiry(token_data) + + if expires_at is None: + return access_token + + now = time.time() + + # Refresh slightly early to avoid the token expiring during the + # PostgreSQL authentication handshake. + if now >= expires_at - _OAUTH_REFRESH_LEEWAY_SECONDS: + logger.info( + "pgAdmin OAuth access token is expired or close to " + "expiration: expires_at=%s now=%s.", + expires_at, + now, + ) + + return _refresh_pgadmin_oauth_token(token_data) + + return access_token + + +def _exchange_oauth_access_token(client_id: str) -> str: + """ + Exchange the current pgAdmin access token for a cluster access token. + + client_id is the cluster's OAuth client identifier, obtained from + the server registration's oauth_client_id connection parameter. + + The token endpoint is obtained from the OAuth provider used for the + current pgAdmin login. An explicitly configured token URL takes + precedence over the token_endpoint discovered through the provider + metadata. + + The current pgAdmin token is refreshed first when necessary. + + Return the exchanged bearer token. Raise OAuthTokenExchangeError + on failure. + + The exchanged token is not stored in the Flask session or cached. + """ + if not has_request_context(): + raise OAuthTokenExchangeError( + "Token exchange requires a pgAdmin login session." + ) + + if ( + not isinstance(client_id, str) or + not client_id.strip() or + "\0" in client_id + ): + raise OAuthTokenExchangeError( + "Set the server connection parameter oauth_client_id " + "to the target cluster identifier." + ) + + subject_token = _get_pgadmin_oauth_token() + + if ( + not isinstance(subject_token, str) or + not subject_token.strip() or + "\0" in subject_token + ): + raise OAuthTokenExchangeError( + "No current pgAdmin OAuth access token is available. " + "Sign in to pgAdmin again." + ) + + try: + oauth_client = _get_current_oauth_client() + except OAuthTokenError as exc: + raise OAuthTokenExchangeError(str(exc)) from None + + try: + metadata = oauth_client.load_server_metadata() + except requests.exceptions.SSLError: + raise OAuthTokenExchangeError( + "TLS verification failed while loading OAuth provider " + "metadata." + ) from None + except requests.exceptions.Timeout: + raise OAuthTokenExchangeError( + "The request for OAuth provider metadata timed out." + ) from None + except requests.exceptions.RequestException: + raise OAuthTokenExchangeError( + "Could not load OAuth provider metadata." + ) from None + except Exception: + raise OAuthTokenExchangeError( + "Could not load OAuth provider metadata." + ) from None + + endpoint = ( + oauth_client.access_token_url or + metadata.get("token_endpoint") + ) + + if not isinstance(endpoint, str) or not endpoint: + raise OAuthTokenExchangeError( + "The OAuth provider has no token endpoint configured " + "or advertised in its metadata." + ) + + try: + parsed_endpoint = urlsplit(endpoint) + valid_endpoint = ( + parsed_endpoint.scheme == "https" and + bool(parsed_endpoint.hostname) and + parsed_endpoint.username is None and + parsed_endpoint.password is None and + not parsed_endpoint.fragment + ) + except ValueError: + valid_endpoint = False + + if not valid_endpoint: + raise OAuthTokenExchangeError( + "The OAuth provider token endpoint must be an HTTPS URL " + "without embedded credentials or a fragment." + ) + + verify_tls = oauth_client.client_kwargs.get("verify", True) + + try: + response = requests.post( + endpoint, + data={ + "client_id": client_id, + "grant_type": ( + "urn:ietf:params:oauth:grant-type:token-exchange" + ), + "subject_token": subject_token, + "subject_token_type": ( + "urn:ietf:params:oauth:token-type:access_token" + ), + "scope": "openid", + }, + headers={ + "Accept": "application/json", + "Content-Type": "application/x-www-form-urlencoded", + }, + auth=_TokenExchangeNoAuth(), + verify=verify_tls, + timeout=(5, 30), + allow_redirects=False, + ) + except requests.exceptions.SSLError: + raise OAuthTokenExchangeError( + "TLS verification failed when contacting the OAuth " + "token endpoint." + ) from None + except requests.exceptions.Timeout: + raise OAuthTokenExchangeError( + "The request to the OAuth token endpoint timed out." + ) from None + except requests.exceptions.RequestException: + raise OAuthTokenExchangeError( + "Could not contact the OAuth token endpoint." + ) from None + + with response: + status_code = response.status_code + + try: + token_data = response.json() + except ValueError: + token_data = None + + has_oauth_error = ( + isinstance(token_data, dict) and + "error" in token_data + ) + + if status_code != 200 or has_oauth_error: + error_code = ( + token_data.get("error") + if isinstance(token_data, dict) + else None + ) + + # Include only a short error identifier. Never expose the + # provider response, descriptions, or token values. + if ( + not isinstance(error_code, str) or + re.fullmatch( + r"[A-Za-z0-9_-]{1,64}", + error_code, + ) is None + ): + error_code = "unknown_error" + + raise OAuthTokenExchangeError( + "OAuth token exchange failed " + "(HTTP {0}, error={1}).".format( + status_code, + error_code, + ) + ) + + if not isinstance(token_data, dict): + raise OAuthTokenExchangeError( + "The OAuth provider returned an invalid token exchange " + "response." + ) + + access_token = token_data.get("access_token") + token_type = token_data.get("token_type") + + if ( + not isinstance(access_token, str) or + not access_token.strip() or + "\0" in access_token + ): + raise OAuthTokenExchangeError( + "The OAuth token exchange response contains no valid " + "access token." + ) + + if ( + not isinstance(token_type, str) or + token_type.lower() != "bearer" + ): + raise OAuthTokenExchangeError( + "The OAuth token exchange response is not a bearer token." + ) + + return access_token + + +def get_postgres_oauth_token( + mode: str, + client_id: Optional[str], +) -> str: + """ + Return a PostgreSQL OAuth token using the selected pgAdmin token mode. + """ + if mode == "direct": + token = _get_pgadmin_oauth_token() + + if not isinstance(token, str) or not token: + raise OAuthTokenError( + "No current pgAdmin OAuth access token is available. " + "Sign in to pgAdmin again." + ) + + return token + + if mode == "exchange": + if ( + not isinstance(client_id, str) or + not client_id.strip() or + "\0" in client_id + ): + raise OAuthTokenError( + "Set the server connection parameter oauth_client_id " + "when using OAuth token exchange." + ) + + return _exchange_oauth_access_token(client_id) + + raise OAuthTokenError( + 'Invalid pgAdmin OAuth token mode "{0}".'.format(mode) + ) diff --git a/web/pgadmin/utils/tests/test_pg_oauth2.py b/web/pgadmin/utils/tests/test_pg_oauth2.py new file mode 100644 index 00000000000..10f90ea264a --- /dev/null +++ b/web/pgadmin/utils/tests/test_pg_oauth2.py @@ -0,0 +1,1273 @@ +########################################################################## +# +# pgAdmin 4 - PostgreSQL Tools +# +# Copyright (C) 2013 - 2026, The pgAdmin Development Team +# This software is released under the PostgreSQL Licence +# +########################################################################## + +"""Tests for pgAdmin PostgreSQL OAuth bearer-token support.""" + +import ctypes +import unittest +from unittest.mock import MagicMock, Mock, patch, create_autospec + +from flask import Flask, session +from pgadmin.utils import pg_oauth2 + +from authlib.integrations.flask_client.apps import FlaskOAuth2App + + +class TestOAuthTokenContext(unittest.TestCase): + """Test propagation of tokens through ContextVar.""" + + def setUp(self): + self.token_handle = pg_oauth2._oauth_token.set(None) + + def tearDown(self): + pg_oauth2._oauth_token.reset(self.token_handle) + + def test_context_sets_and_restores_token(self): + self.assertIsNone(pg_oauth2._oauth_token.get()) + + with pg_oauth2.oauth_token_context('access-token'): + self.assertEqual( + pg_oauth2._oauth_token.get(), + 'access-token' + ) + + self.assertIsNone(pg_oauth2._oauth_token.get()) + + def test_context_restores_existing_token(self): + existing_handle = pg_oauth2._oauth_token.set( + 'existing-token' + ) + + try: + with pg_oauth2.oauth_token_context('temporary-token'): + self.assertEqual( + pg_oauth2._oauth_token.get(), + 'temporary-token' + ) + + self.assertEqual( + pg_oauth2._oauth_token.get(), + 'existing-token' + ) + finally: + pg_oauth2._oauth_token.reset(existing_handle) + + def test_nested_contexts_restore_previous_token(self): + with pg_oauth2.oauth_token_context('outer-token'): + self.assertEqual( + pg_oauth2._oauth_token.get(), + 'outer-token' + ) + + with pg_oauth2.oauth_token_context('inner-token'): + self.assertEqual( + pg_oauth2._oauth_token.get(), + 'inner-token' + ) + + self.assertEqual( + pg_oauth2._oauth_token.get(), + 'outer-token' + ) + + self.assertIsNone(pg_oauth2._oauth_token.get()) + + def test_context_restores_token_when_exception_is_raised(self): + with self.assertRaisesRegex(RuntimeError, 'test error'): + with pg_oauth2.oauth_token_context('access-token'): + raise RuntimeError('test error') + + self.assertIsNone(pg_oauth2._oauth_token.get()) + + +class TestOAuthHook(unittest.TestCase): + """Test the ctypes libpq OAuth callback.""" + + def setUp(self): + self.token_handle = pg_oauth2._oauth_token.set(None) + + self.previous_hook = pg_oauth2._previous_hook + pg_oauth2._previous_hook = None + + self.previous_token_buffers = dict( + pg_oauth2._token_buffers + ) + pg_oauth2._token_buffers.clear() + + def tearDown(self): + pg_oauth2._token_buffers.clear() + pg_oauth2._token_buffers.update( + self.previous_token_buffers + ) + + pg_oauth2._previous_hook = self.previous_hook + pg_oauth2._oauth_token.reset(self.token_handle) + + @staticmethod + def _make_request(): + request = pg_oauth2.PGoauthBearerRequest() + request_pointer = ctypes.pointer(request) + data = ctypes.cast( + request_pointer, + ctypes.c_void_p + ) + + return request, request_pointer, data + + def test_non_bearer_request_delegates_to_previous_hook(self): + previous_hook = Mock(return_value=73) + pg_oauth2._previous_hook = previous_hook + + result = pg_oauth2._oauth_hook( + 0, + None, + None + ) + + self.assertEqual(result, 73) + previous_hook.assert_called_once_with( + 0, + None, + None + ) + + def test_non_bearer_request_returns_zero_without_previous_hook(self): + result = pg_oauth2._oauth_hook( + 0, + None, + None + ) + + self.assertEqual(result, 0) + + def test_missing_token_delegates_to_previous_hook(self): + previous_hook = Mock(return_value=41) + pg_oauth2._previous_hook = previous_hook + + result = pg_oauth2._oauth_hook( + pg_oauth2.PQAUTHDATA_OAUTH_BEARER_TOKEN, + None, + None + ) + + self.assertEqual(result, 41) + previous_hook.assert_called_once_with( + pg_oauth2.PQAUTHDATA_OAUTH_BEARER_TOKEN, + None, + None + ) + + def test_missing_token_returns_zero_without_previous_hook(self): + result = pg_oauth2._oauth_hook( + pg_oauth2.PQAUTHDATA_OAUTH_BEARER_TOKEN, + None, + None + ) + + self.assertEqual(result, 0) + + def test_bearer_request_receives_token(self): + previous_hook = Mock(return_value=51) + pg_oauth2._previous_hook = previous_hook + + request, request_pointer, data = self._make_request() + + with pg_oauth2.oauth_token_context('test-access-token'): + result = pg_oauth2._oauth_hook( + pg_oauth2.PQAUTHDATA_OAUTH_BEARER_TOKEN, + None, + data + ) + + self.assertEqual(result, 1) + self.assertEqual(request.token, b'test-access-token') + self.assertIsNone(request.async_) + self.assertIsNotNone(request.cleanup) + + request_address = ctypes.addressof(request) + self.assertIn(request_address, pg_oauth2._token_buffers) + + previous_hook.assert_not_called() + + # Keep the ctypes pointer alive for the whole test. + self.assertIsNotNone(request_pointer) + + def test_cleanup_removes_retained_token_buffer(self): + request, request_pointer, data = self._make_request() + + with pg_oauth2.oauth_token_context( + 'test-access-token'): + result = pg_oauth2._oauth_hook( + pg_oauth2.PQAUTHDATA_OAUTH_BEARER_TOKEN, + None, + data + ) + + self.assertEqual(result, 1) + + request_address = ctypes.addressof(request) + + self.assertIn( + request_address, + pg_oauth2._token_buffers + ) + + pg_oauth2._oauth_cleanup(None, data) + + self.assertNotIn( + request_address, + pg_oauth2._token_buffers + ) + + # Keep request_pointer alive until cleanup has completed. + self.assertIsNotNone(request_pointer) + + def test_previous_hook_exception_returns_minus_one(self): + """ + Never allow an exception from the previous libpq hook to cross + the native callback boundary. + """ + previous_hook = Mock( + side_effect=RuntimeError('previous hook failed') + ) + pg_oauth2._previous_hook = previous_hook + + result = pg_oauth2._oauth_hook( + 0, + None, + None + ) + + self.assertEqual(result, -1) + previous_hook.assert_called_once_with( + 0, + None, + None + ) + + def test_bearer_hook_internal_failure_returns_minus_one(self): + """ + Never allow an exception raised while populating the bearer + request to cross the native callback boundary. + """ + previous_hook = Mock(return_value=51) + pg_oauth2._previous_hook = previous_hook + + request, request_pointer, data = self._make_request() + + # An access token must be a string. This deliberately supplies a + # malformed value so token.encode() raises inside the callback. + with pg_oauth2.oauth_token_context(123): + result = pg_oauth2._oauth_hook( + pg_oauth2.PQAUTHDATA_OAUTH_BEARER_TOKEN, + None, + data + ) + + self.assertEqual(result, -1) + self.assertIsNone(request.token) + self.assertNotIn( + ctypes.addressof(request), + pg_oauth2._token_buffers + ) + + # Once pgAdmin chooses to handle a bearer request, an internal + # failure must not invoke the previous hook afterward. + previous_hook.assert_not_called() + + # Keep the ctypes allocation alive for the complete callback. + self.assertIsNotNone(request_pointer) + + +class TestGetCurrentOAuthClient(unittest.TestCase): + """Test lookup of the Authlib client for the current OAuth provider.""" + + def setUp(self): + self.app = Flask(__name__) + self.app.config.update( + SECRET_KEY='pg-oauth2-unit-test-secret', + TESTING=True, + ) + + @patch('pgadmin.authenticate.get_auth_sources') + def test_unknown_provider_raises_error(self, get_auth_sources): + oauth_source = MagicMock() + oauth_source.oauth2_clients = {} + get_auth_sources.return_value = oauth_source + + with self.app.test_request_context('/'): + session['oauth2_provider'] = 'unknown-provider' + + with self.assertRaises(pg_oauth2.OAuthTokenError): + pg_oauth2._get_current_oauth_client() + + @patch('pgadmin.authenticate.get_auth_sources') + def test_returns_registered_oauth_client(self, get_auth_sources): + oauth_client = MagicMock() + oauth_source = MagicMock() + oauth_source.oauth2_clients = { + 'keycloak': oauth_client, + } + get_auth_sources.return_value = oauth_source + + with self.app.test_request_context('/'): + session['oauth2_provider'] = 'keycloak' + + result = pg_oauth2._get_current_oauth_client() + + self.assertIs(result, oauth_client) + + +class TestGetPgAdminOAuthToken(unittest.TestCase): + """Test pgAdmin access-token retrieval and refresh.""" + + PROVIDER_NAME = 'keycloak' + + def setUp(self): + self.app = Flask(__name__) + self.app.config.update( + SECRET_KEY='pg-oauth2-unit-test-secret', + TESTING=True, + ) + + def _make_oauth_client(self): + return create_autospec( + FlaskOAuth2App, + instance=True, + spec_set=True + ) + + def test_returns_none_without_request_context(self): + self.assertIsNone( + pg_oauth2._get_pgadmin_oauth_token() + ) + + def test_returns_none_without_oauth_token(self): + with self.app.test_request_context('/'): + self.assertIsNone( + pg_oauth2._get_pgadmin_oauth_token() + ) + + def test_returns_none_when_token_data_is_not_dictionary(self): + with self.app.test_request_context('/'): + session['oauth2_token'] = 'not-a-dictionary' + + self.assertIsNone( + pg_oauth2._get_pgadmin_oauth_token() + ) + + def test_returns_none_when_access_token_is_missing(self): + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'refresh_token': 'refresh-token', + } + + self.assertIsNone( + pg_oauth2._get_pgadmin_oauth_token() + ) + + def test_returns_current_unexpired_access_token(self): + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'current-access-token', + 'refresh_token': 'refresh-token', + 'expires_at': 2000, + } + + with patch( + 'pgadmin.utils.pg_oauth2.time.time', + return_value=1000 + ): + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertEqual( + result, + 'current-access-token' + ) + + @patch( + 'pgadmin.utils.pg_oauth2._get_current_oauth_client' + ) + def test_refreshes_expired_access_token(self, get_oauth_client): + oauth_client = self._make_oauth_client() + oauth_client.fetch_access_token.return_value = { + 'access_token': 'refreshed-access-token', + 'refresh_token': 'new-refresh-token', + } + + get_oauth_client.return_value = oauth_client + + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'expired-access-token', + 'refresh_token': 'original-refresh-token', + 'expires_at': 999, + 'token_type': 'Bearer', + } + session['oauth2_provider'] = self.PROVIDER_NAME + + with patch( + 'pgadmin.utils.pg_oauth2.time.time', + return_value=1000 + ): + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertEqual(result, 'refreshed-access-token') + + oauth_client.fetch_access_token.assert_called_once_with( + grant_type='refresh_token', + refresh_token='original-refresh-token' + ) + + @patch( + 'pgadmin.utils.pg_oauth2._get_current_oauth_client' + ) + def test_refresh_failure_returns_none(self, get_oauth_client): + oauth_client = self._make_oauth_client() + oauth_client.fetch_access_token.side_effect = RuntimeError( + 'token endpoint unavailable' + ) + get_oauth_client.return_value = oauth_client + + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'expired-access-token', + 'refresh_token': 'refresh-token', + 'expires_at': 999, + 'token_type': 'Bearer', + } + session['oauth2_provider'] = self.PROVIDER_NAME + + with patch( + 'pgadmin.utils.pg_oauth2.time.time', + return_value=1000 + ): + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertIsNone(result) + oauth_client.fetch_access_token.assert_called_once_with( + grant_type='refresh_token', + refresh_token='refresh-token' + ) + + @patch( + 'pgadmin.utils.pg_oauth2._get_current_oauth_client' + ) + def test_expired_token_without_refresh_token_returns_none( + self, get_oauth_client): + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'expired-access-token', + 'expires_at': 999, + 'token_type': 'Bearer', + } + session['oauth2_provider'] = self.PROVIDER_NAME + + with patch( + 'pgadmin.utils.pg_oauth2.time.time', + return_value=1000 + ): + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertIsNone(result) + get_oauth_client.assert_not_called() + + @patch( + 'pgadmin.utils.pg_oauth2._get_current_oauth_client' + ) + def test_refreshes_access_token_within_expiry_safety_window( + self, get_oauth_client): + """ + Refresh a token that has not expired yet but will expire within + the 30-second safety window. + """ + oauth_client = self._make_oauth_client() + oauth_client.fetch_access_token.return_value = { + 'access_token': 'refreshed-access-token', + 'refresh_token': 'refreshed-refresh-token', + 'expires_at': 3000, + 'token_type': 'Bearer', + } + get_oauth_client.return_value = oauth_client + + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'nearly-expired-access-token', + 'refresh_token': 'original-refresh-token', + 'expires_at': 1020, + 'token_type': 'Bearer', + } + session['oauth2_provider'] = self.PROVIDER_NAME + + with patch( + 'pgadmin.utils.pg_oauth2.time.time', + return_value=1000 + ): + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertEqual(result, 'refreshed-access-token') + self.assertEqual( + session['oauth2_token']['access_token'], + 'refreshed-access-token' + ) + + oauth_client.fetch_access_token.assert_called_once_with( + grant_type='refresh_token', + refresh_token='original-refresh-token' + ) + + def test_returns_access_token_without_expiration(self): + """ + Return an access token when the provider did not supply expires_at. + """ + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'access-token-without-expiration', + 'refresh_token': 'refresh-token', + 'token_type': 'Bearer', + } + + with patch( + 'pgadmin.utils.pg_oauth2.' + '_refresh_pgadmin_oauth_token' + ) as refresh_mock: + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertEqual( + result, + 'access-token-without-expiration' + ) + refresh_mock.assert_not_called() + + @patch( + 'pgadmin.utils.pg_oauth2._get_current_oauth_client' + ) + def test_refresh_preserves_existing_refresh_token( + self, get_oauth_client): + """ + Preserve the old refresh token when the provider returns only a + new access token. + """ + oauth_client = self._make_oauth_client() + oauth_client.fetch_access_token.return_value = { + 'access_token': 'refreshed-access-token', + 'expires_at': 3000, + 'token_type': 'Bearer', + } + get_oauth_client.return_value = oauth_client + + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'expired-access-token', + 'refresh_token': 'original-refresh-token', + 'expires_at': 900, + 'token_type': 'Bearer', + } + session['oauth2_provider'] = self.PROVIDER_NAME + + with patch( + 'pgadmin.utils.pg_oauth2.time.time', + return_value=1000 + ): + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertEqual(result, 'refreshed-access-token') + self.assertEqual( + session['oauth2_token']['refresh_token'], + 'original-refresh-token' + ) + + oauth_client.fetch_access_token.assert_called_once_with( + grant_type='refresh_token', + refresh_token='original-refresh-token' + ) + + def test_expired_token_without_provider_returns_none(self): + """ + An expired token cannot be refreshed when the OAuth provider + identity is missing from the session. + """ + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'expired-access-token', + 'refresh_token': 'refresh-token', + 'expires_at': 900, + 'token_type': 'Bearer', + } + + with patch( + 'pgadmin.utils.pg_oauth2.time.time', + return_value=1000 + ): + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertIsNone(result) + + @patch( + 'pgadmin.utils.pg_oauth2._get_current_oauth_client', + side_effect=pg_oauth2.OAuthTokenError( + 'Unknown OAuth provider.' + ) + ) + def test_expired_token_returns_none_when_client_lookup_fails( + self, get_oauth_client): + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'expired-access-token', + 'refresh_token': 'refresh-token', + 'expires_at': 900, + 'token_type': 'Bearer', + } + session['oauth2_provider'] = 'unknown-provider' + + with patch( + 'pgadmin.utils.pg_oauth2.time.time', + return_value=1000 + ): + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertIsNone(result) + get_oauth_client.assert_called_once_with() + + @patch( + 'pgadmin.utils.pg_oauth2._get_current_oauth_client' + ) + def test_refresh_returns_none_for_non_dictionary_response( + self, get_oauth_client): + """ + Reject a malformed refresh response instead of attempting to use + or store it as an OAuth token. + """ + oauth_client = self._make_oauth_client() + oauth_client.fetch_access_token.return_value = ( + 'invalid-token-response' + ) + get_oauth_client.return_value = oauth_client + + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'expired-access-token', + 'refresh_token': 'original-refresh-token', + 'expires_at': 900, + 'token_type': 'Bearer', + } + session['oauth2_provider'] = self.PROVIDER_NAME + + with patch( + 'pgadmin.utils.pg_oauth2.time.time', + return_value=1000 + ): + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertIsNone(result) + self.assertEqual( + session['oauth2_token']['access_token'], + 'expired-access-token' + ) + + oauth_client.fetch_access_token.assert_called_once_with( + grant_type='refresh_token', + refresh_token='original-refresh-token' + ) + + @patch( + 'pgadmin.utils.pg_oauth2._get_current_oauth_client' + ) + def test_refresh_returns_none_when_access_token_is_missing( + self, get_oauth_client): + """ + Reject a syntactically valid refresh response that does not contain + a new access token. + """ + oauth_client = self._make_oauth_client() + oauth_client.fetch_access_token.return_value = { + 'refresh_token': 'new-refresh-token', + 'expires_at': 3000, + 'token_type': 'Bearer', + } + get_oauth_client.return_value = oauth_client + + with self.app.test_request_context('/'): + session['oauth2_token'] = { + 'access_token': 'expired-access-token', + 'refresh_token': 'original-refresh-token', + 'expires_at': 900, + 'token_type': 'Bearer', + } + session['oauth2_provider'] = self.PROVIDER_NAME + + with patch( + 'pgadmin.utils.pg_oauth2.time.time', + return_value=1000 + ): + result = pg_oauth2._get_pgadmin_oauth_token() + + self.assertIsNone(result) + self.assertEqual( + session['oauth2_token']['access_token'], + 'expired-access-token' + ) + + oauth_client.fetch_access_token.assert_called_once_with( + grant_type='refresh_token', + refresh_token='original-refresh-token' + ) + +class TestExchangeOAuthAccessToken(unittest.TestCase): + """Test OAuth access-token exchange for PostgreSQL.""" + + TOKEN_ENDPOINT = ( + 'https://keycloak.example.test/realms/postgres/' + 'protocol/openid-connect/token' + ) + + DISCOVERED_TOKEN_ENDPOINT = ( + 'https://keycloak.example.test/discovered/token' + ) + + def setUp(self): + self.app = Flask(__name__) + self.app.config.update( + SECRET_KEY='pg-oauth2-unit-test-secret', + TESTING=True, + ) + + @staticmethod + def _make_oauth_client( + access_token_url=None, + verify=True): + oauth_client = MagicMock() + oauth_client.access_token_url = access_token_url + oauth_client.client_kwargs = { + 'verify': verify, + } + oauth_client.load_server_metadata.return_value = { + 'token_endpoint': ( + TestExchangeOAuthAccessToken.DISCOVERED_TOKEN_ENDPOINT + ), + } + + return oauth_client + + @staticmethod + def _make_response(status_code=200, token_data=None): + response = MagicMock() + response.status_code = status_code + response.json.return_value = token_data + response.__enter__.return_value = response + response.__exit__.return_value = False + + return response + + def test_exchange_requires_request_context(self): + with self.assertRaisesRegex( + pg_oauth2.OAuthTokenExchangeError, + 'pgAdmin login session'): + pg_oauth2._exchange_oauth_access_token( + 'postgres-cluster' + ) + + @patch('pgadmin.utils.pg_oauth2.requests.post') + @patch('pgadmin.utils.pg_oauth2._get_current_oauth_client') + @patch('pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token') + def test_exchange_returns_access_token( + self, get_token, get_oauth_client, post): + get_token.return_value = 'subject-access-token' + + oauth_client = self._make_oauth_client( + access_token_url=self.TOKEN_ENDPOINT + ) + get_oauth_client.return_value = oauth_client + + post.return_value = self._make_response( + token_data={ + 'access_token': 'exchanged-access-token', + 'token_type': 'Bearer', + } + ) + + with self.app.test_request_context('/'): + result = pg_oauth2._exchange_oauth_access_token( + 'postgres-cluster' + ) + + self.assertEqual(result, 'exchanged-access-token') + + get_token.assert_called_once_with() + get_oauth_client.assert_called_once_with() + post.assert_called_once() + + args, kwargs = post.call_args + + self.assertEqual(args[0], self.TOKEN_ENDPOINT) + + self.assertEqual( + kwargs['data'], + { + 'client_id': 'postgres-cluster', + 'grant_type': ( + 'urn:ietf:params:oauth:grant-type:token-exchange' + ), + 'subject_token': 'subject-access-token', + 'subject_token_type': ( + 'urn:ietf:params:oauth:token-type:access_token' + ), + 'scope': 'openid', + } + ) + + self.assertEqual( + kwargs['headers'], + { + 'Accept': 'application/json', + 'Content-Type': 'application/x-www-form-urlencoded', + } + ) + + self.assertIsInstance( + kwargs['auth'], + pg_oauth2._TokenExchangeNoAuth + ) + self.assertTrue(kwargs['verify']) + self.assertEqual(kwargs['timeout'], (5, 30)) + self.assertFalse(kwargs['allow_redirects']) + + @patch('pgadmin.utils.pg_oauth2.requests.post') + @patch('pgadmin.utils.pg_oauth2._get_current_oauth_client') + @patch('pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token') + def test_exchange_uses_discovered_token_endpoint( + self, get_token, get_oauth_client, post): + get_token.return_value = 'subject-access-token' + + oauth_client = self._make_oauth_client( + access_token_url=None + ) + get_oauth_client.return_value = oauth_client + + post.return_value = self._make_response( + token_data={ + 'access_token': 'exchanged-access-token', + 'token_type': 'Bearer', + } + ) + + with self.app.test_request_context('/'): + result = pg_oauth2._exchange_oauth_access_token( + 'postgres-cluster' + ) + + self.assertEqual(result, 'exchanged-access-token') + + oauth_client.load_server_metadata.assert_called_once_with() + post.assert_called_once() + + args, _ = post.call_args + self.assertEqual( + args[0], + self.DISCOVERED_TOKEN_ENDPOINT + ) + + @patch('pgadmin.utils.pg_oauth2.requests.post') + @patch('pgadmin.utils.pg_oauth2._get_current_oauth_client') + @patch('pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token') + def test_exchange_fails_when_provider_metadata_cannot_be_loaded( + self, get_token, get_oauth_client, post): + get_token.return_value = 'subject-access-token' + + oauth_client = self._make_oauth_client( + access_token_url=None + ) + oauth_client.load_server_metadata.side_effect = RuntimeError( + 'metadata unavailable' + ) + get_oauth_client.return_value = oauth_client + + with self.app.test_request_context('/'): + with self.assertRaises( + pg_oauth2.OAuthTokenExchangeError): + pg_oauth2._exchange_oauth_access_token( + 'postgres-cluster' + ) + + oauth_client.load_server_metadata.assert_called_once_with() + post.assert_not_called() + + @patch('pgadmin.utils.pg_oauth2.requests.post') + @patch('pgadmin.utils.pg_oauth2._get_current_oauth_client') + @patch('pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token') + def test_exchange_uses_provider_tls_verification_setting( + self, get_token, get_oauth_client, post): + get_token.return_value = 'subject-access-token' + + oauth_client = self._make_oauth_client( + access_token_url=self.TOKEN_ENDPOINT, + verify=False + ) + get_oauth_client.return_value = oauth_client + + post.return_value = self._make_response( + token_data={ + 'access_token': 'exchanged-access-token', + 'token_type': 'Bearer', + } + ) + + with self.app.test_request_context('/'): + result = pg_oauth2._exchange_oauth_access_token( + 'postgres-cluster' + ) + + self.assertEqual(result, 'exchanged-access-token') + + post.assert_called_once() + _, kwargs = post.call_args + self.assertFalse(kwargs['verify']) + + @patch('pgadmin.utils.pg_oauth2.requests.post') + @patch('pgadmin.utils.pg_oauth2._get_current_oauth_client') + @patch('pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token') + def test_exchange_translates_request_timeout( + self, get_token, get_oauth_client, post): + get_token.return_value = 'subject-access-token' + + get_oauth_client.return_value = self._make_oauth_client( + access_token_url=self.TOKEN_ENDPOINT + ) + + post.side_effect = pg_oauth2.requests.exceptions.Timeout( + 'request timed out' + ) + + with self.app.test_request_context('/'): + with self.assertRaisesRegex( + pg_oauth2.OAuthTokenExchangeError, + 'timed out'): + pg_oauth2._exchange_oauth_access_token( + 'postgres-cluster' + ) + + post.assert_called_once() + + @patch('pgadmin.utils.pg_oauth2.requests.post') + @patch('pgadmin.utils.pg_oauth2._get_current_oauth_client') + @patch('pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token') + def test_exchange_rejects_oauth_error_response( + self, get_token, get_oauth_client, post): + get_token.return_value = 'subject-access-token' + + get_oauth_client.return_value = self._make_oauth_client( + access_token_url=self.TOKEN_ENDPOINT + ) + + post.return_value = self._make_response( + status_code=400, + token_data={ + 'error': 'invalid_grant', + } + ) + + with self.app.test_request_context('/'): + with self.assertRaisesRegex( + pg_oauth2.OAuthTokenExchangeError, + 'invalid_grant'): + pg_oauth2._exchange_oauth_access_token( + 'postgres-cluster' + ) + + post.assert_called_once() + + @patch('pgadmin.utils.pg_oauth2.requests.post') + @patch('pgadmin.utils.pg_oauth2._get_current_oauth_client') + @patch('pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token') + def test_exchange_rejects_non_dictionary_response( + self, get_token, get_oauth_client, post): + get_token.return_value = 'subject-access-token' + + get_oauth_client.return_value = self._make_oauth_client( + access_token_url=self.TOKEN_ENDPOINT + ) + + post.return_value = self._make_response( + token_data='invalid-response' + ) + + with self.app.test_request_context('/'): + with self.assertRaises( + pg_oauth2.OAuthTokenExchangeError): + pg_oauth2._exchange_oauth_access_token( + 'postgres-cluster' + ) + + post.assert_called_once() + + @patch('pgadmin.utils.pg_oauth2.requests.post') + @patch('pgadmin.utils.pg_oauth2._get_current_oauth_client') + @patch('pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token') + def test_exchange_rejects_invalid_access_token( + self, get_token, get_oauth_client, post): + get_token.return_value = 'subject-access-token' + + get_oauth_client.return_value = self._make_oauth_client( + access_token_url=self.TOKEN_ENDPOINT + ) + + invalid_tokens = ( + None, + '', + ' ', + 'invalid\0token', + ) + + with self.app.test_request_context('/'): + for access_token in invalid_tokens: + with self.subTest(access_token=access_token): + post.return_value = self._make_response( + token_data={ + 'access_token': access_token, + 'token_type': 'Bearer', + } + ) + + with self.assertRaises( + pg_oauth2.OAuthTokenExchangeError): + pg_oauth2._exchange_oauth_access_token( + 'postgres-cluster' + ) + + self.assertEqual( + post.call_count, + len(invalid_tokens) + ) + + @patch('pgadmin.utils.pg_oauth2.requests.post') + @patch('pgadmin.utils.pg_oauth2._get_current_oauth_client') + @patch('pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token') + def test_exchange_rejects_non_bearer_token( + self, get_token, get_oauth_client, post): + get_token.return_value = 'subject-access-token' + + get_oauth_client.return_value = self._make_oauth_client( + access_token_url=self.TOKEN_ENDPOINT + ) + + post.return_value = self._make_response( + token_data={ + 'access_token': 'exchanged-access-token', + 'token_type': 'DPoP', + } + ) + + with self.app.test_request_context('/'): + with self.assertRaises( + pg_oauth2.OAuthTokenExchangeError): + pg_oauth2._exchange_oauth_access_token( + 'postgres-cluster' + ) + + post.assert_called_once() + + +class TestGetPostgresOAuthToken(unittest.TestCase): + """Test selection of the PostgreSQL OAuth token source.""" + + @patch( + 'pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token', + return_value='pgadmin-access-token' + ) + def test_direct_mode_returns_pgadmin_access_token(self, get_token): + result = pg_oauth2.get_postgres_oauth_token( + 'direct', + None + ) + + self.assertEqual(result, 'pgadmin-access-token') + get_token.assert_called_once_with() + + @patch( + 'pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token', + return_value=None + ) + def test_direct_mode_raises_when_pgadmin_token_is_missing( + self, get_token): + with self.assertRaisesRegex( + pg_oauth2.OAuthTokenError, + 'No current pgAdmin OAuth access token is available'): + pg_oauth2.get_postgres_oauth_token( + 'direct', + None + ) + + get_token.assert_called_once_with() + + @patch( + 'pgadmin.utils.pg_oauth2._exchange_oauth_access_token', + return_value='exchanged-access-token' + ) + def test_exchange_mode_returns_exchanged_access_token(self, exchange): + result = pg_oauth2.get_postgres_oauth_token( + 'exchange', + 'postgres-cluster' + ) + + self.assertEqual(result, 'exchanged-access-token') + exchange.assert_called_once_with('postgres-cluster') + + @patch('pgadmin.utils.pg_oauth2._exchange_oauth_access_token') + def test_exchange_mode_requires_client_id(self, exchange): + invalid_client_ids = ( + None, + '', + ' ', + 'cluster\0name', + ) + + for client_id in invalid_client_ids: + with self.subTest(client_id=client_id): + with self.assertRaisesRegex( + pg_oauth2.OAuthTokenError, + 'oauth_client_id'): + pg_oauth2.get_postgres_oauth_token( + 'exchange', + client_id + ) + + exchange.assert_not_called() + + @patch('pgadmin.utils.pg_oauth2._get_pgadmin_oauth_token') + @patch('pgadmin.utils.pg_oauth2._exchange_oauth_access_token') + def test_invalid_mode_raises_error(self, exchange, get_token): + with self.assertRaisesRegex( + pg_oauth2.OAuthTokenError, + 'Invalid pgAdmin OAuth token mode'): + pg_oauth2.get_postgres_oauth_token( + 'invalid-mode', + 'postgres-cluster' + ) + + get_token.assert_not_called() + exchange.assert_not_called() + + +class TestInstallOAuthHook(unittest.TestCase): + """Test installation of the process-global libpq OAuth hook.""" + + def setUp(self): + self.original_libpq = pg_oauth2._libpq + self.original_hook = pg_oauth2._hook + self.original_cleanup_hook = pg_oauth2._cleanup_hook + self.original_previous_hook = pg_oauth2._previous_hook + + pg_oauth2._libpq = None + pg_oauth2._hook = None + pg_oauth2._cleanup_hook = None + pg_oauth2._previous_hook = None + + def tearDown(self): + pg_oauth2._libpq = self.original_libpq + pg_oauth2._hook = self.original_hook + pg_oauth2._cleanup_hook = self.original_cleanup_hook + pg_oauth2._previous_hook = self.original_previous_hook + + @staticmethod + def _make_libpq(previous_hook_pointer=0): + libpq = MagicMock(spec_set=[ + 'PQlibVersion', + 'PQgetAuthDataHook', + 'PQsetAuthDataHook', + ]) + + libpq.PQlibVersion.return_value = 180000 + + installed_hook_pointer = ctypes.cast( + pg_oauth2._oauth_hook, + ctypes.c_void_p + ).value + + # The first call obtains the previous hook. Some versions of the + # implementation make a second call to verify installation. + libpq.PQgetAuthDataHook.side_effect = [ + previous_hook_pointer, + installed_hook_pointer, + ] + + return libpq + + def test_installs_oauth_hook(self): + libpq = self._make_libpq() + + with patch( + 'pgadmin.utils.pg_oauth2._get_libpq', + return_value=libpq) as get_libpq: + pg_oauth2.install_oauth_hook() + + get_libpq.assert_called_once_with() + libpq.PQsetAuthDataHook.assert_called_once_with( + pg_oauth2._oauth_hook + ) + + self.assertIs( + pg_oauth2._hook, + pg_oauth2._oauth_hook + ) + self.assertIs( + pg_oauth2._cleanup_hook, + pg_oauth2._oauth_cleanup + ) + self.assertIsNone(pg_oauth2._previous_hook) + + def test_installation_is_idempotent(self): + libpq = self._make_libpq() + + with patch( + 'pgadmin.utils.pg_oauth2._get_libpq', + return_value=libpq) as get_libpq: + pg_oauth2.install_oauth_hook() + pg_oauth2.install_oauth_hook() + + # The second call should return before loading libpq or replacing + # the process-global hook again. + get_libpq.assert_called_once_with() + libpq.PQsetAuthDataHook.assert_called_once_with( + pg_oauth2._oauth_hook + ) + + def test_preserves_previous_auth_data_hook(self): + @pg_oauth2._AUTH_HOOK + def previous_hook(authdata_type, conn, data): + return 67 + + previous_hook_pointer = ctypes.cast( + previous_hook, + ctypes.c_void_p + ).value + + libpq = self._make_libpq(previous_hook_pointer) + + with patch( + 'pgadmin.utils.pg_oauth2._get_libpq', + return_value=libpq): + pg_oauth2.install_oauth_hook() + + self.assertIsNotNone(pg_oauth2._previous_hook) + self.assertEqual( + pg_oauth2._previous_hook(0, None, None), + 67 + ) + + libpq.PQsetAuthDataHook.assert_called_once_with( + pg_oauth2._oauth_hook + ) + + # Keep the original ctypes callback alive until after the + # reconstructed callback pointer has been invoked. + self.assertIsNotNone(previous_hook) + + +if __name__ == '__main__': + unittest.main()