diff --git a/google/cloud/sql/connector/connection_info.py b/google/cloud/sql/connector/connection_info.py index 49404f49..01033863 100644 --- a/google/cloud/sql/connector/connection_info.py +++ b/google/cloud/sql/connector/connection_info.py @@ -45,6 +45,10 @@ async def connect_info(self) -> ConnectionInfo: async def force_refresh(self) -> None: pass + @abc.abstractmethod + def record_principal(self, user: str, database: str) -> None: + pass + @abc.abstractmethod async def close(self) -> None: pass diff --git a/google/cloud/sql/connector/connector.py b/google/cloud/sql/connector/connector.py index 7a9964ca..73c0b04a 100644 --- a/google/cloud/sql/connector/connector.py +++ b/google/cloud/sql/connector/connector.py @@ -338,6 +338,8 @@ async def connect_async( self._client, self._keys, enable_iam_auth, + ip_type=self._ip_type, + timeout=self._timeout, ) else: logger.debug( @@ -348,6 +350,8 @@ async def connect_async( self._client, self._keys, enable_iam_auth, + ip_type=self._ip_type, + timeout=self._timeout, ) # wrap cache as a MonitoredCache monitored_cache = MonitoredCache( @@ -437,6 +441,9 @@ async def connect_async( f"['{instance_connection_string}']: Truncated IAM database username from {kwargs['user']} to {formatted_user}" ) kwargs["user"] = formatted_user + if conn_info.database_version.startswith("POSTGRES"): + db_name = kwargs.get("db") or kwargs.get("database") or kwargs["user"] + monitored_cache.record_principal(kwargs["user"], str(db_name)) try: last_ex = None diff --git a/google/cloud/sql/connector/instance.py b/google/cloud/sql/connector/instance.py index 28ab54e4..eeb297b2 100644 --- a/google/cloud/sql/connector/instance.py +++ b/google/cloud/sql/connector/instance.py @@ -26,6 +26,7 @@ from google.cloud.sql.connector.connection_info import ConnectionInfo from google.cloud.sql.connector.connection_info import ConnectionInfoCache from google.cloud.sql.connector.connection_name import ConnectionName +from google.cloud.sql.connector.enums import IPTypes from google.cloud.sql.connector.exceptions import RefreshNotValidError from google.cloud.sql.connector.rate_limiter import AsyncRateLimiter from google.cloud.sql.connector.refresh_utils import _is_valid @@ -34,6 +35,8 @@ logger = logging.getLogger(name=__name__) APPLICATION_NAME = "cloud-sql-python-connector" +SERVER_PROXY_PORT = 3307 +DEFAULT_CONNECT_TIMEOUT = 30 class RefreshAheadCache(ConnectionInfoCache): @@ -50,6 +53,8 @@ def __init__( client: CloudSQLClient, keys: asyncio.Future, enable_iam_auth: bool = False, + ip_type: IPTypes | str = IPTypes.PUBLIC, + timeout: int = DEFAULT_CONNECT_TIMEOUT, ) -> None: """Initializes a RefreshAheadCache instance. @@ -62,10 +67,16 @@ def __init__( enable_iam_auth (bool): Enables automatic IAM database authentication (Postgres and MySQL) as the default authentication method for all connections. + ip_type (IPTypes | str): Preferred IP type used to connect to the instance. + timeout (int): Connect timeout in seconds. """ self._conn_name = conn_name self._enable_iam_auth = enable_iam_auth + if isinstance(ip_type, str): + ip_type = IPTypes._from_str(ip_type) + self._ip_type = ip_type + self._timeout = timeout self._keys = keys self._client = client self._refresh_rate_limiter = AsyncRateLimiter( @@ -73,10 +84,15 @@ def __init__( rate=1 / 30, ) self._refresh_in_progress = asyncio.locks.Event() + self._iam_principals: list[tuple[str, str]] = [] self._current: asyncio.Task = self._schedule_refresh(0) self._next: asyncio.Task = self._current self._closed = False + def record_principal(self, user: str, database: str) -> None: + """Records a PostgreSQL (user, database) pair observed on an IAM connection.""" + _append_iam_principal(self._iam_principals, user, database) + @property def conn_name(self) -> ConnectionName: return self._conn_name @@ -119,6 +135,8 @@ async def _perform_refresh(self) -> ConnectionInfo: self._keys, self._enable_iam_auth, ) + if self._enable_iam_auth: + await self._probe_connection(connection_info) logger.debug( f"['{self._conn_name}']: Connection info refresh operation complete" ) @@ -138,6 +156,17 @@ async def _perform_refresh(self) -> ConnectionInfo: self._refresh_in_progress.clear() return connection_info + async def _probe_connection(self, conn_info: ConnectionInfo) -> None: + """Proactively probes the database to refresh IAM tokens on server-side MCP.""" + await _probe_instance_connection( + self._conn_name, + conn_info, + self._enable_iam_auth, + self._ip_type, + self._timeout, + self._iam_principals, + ) + def _schedule_refresh(self, delay: int) -> asyncio.Task: """ Schedule task to sleep and then perform refresh to get ConnectionInfo. @@ -223,3 +252,124 @@ async def close(self) -> None: tasks = asyncio.gather(self._current, self._next, return_exceptions=True) await asyncio.wait_for(tasks, timeout=2.0) self._closed = True + + +def _append_iam_principal( + principals: list[tuple[str, str]], user: str, database: str +) -> None: + """Appends (user, database) to principals if not already present.""" + if not user: + return + if not database: + database = user + pair = (user, database) + if pair not in principals: + principals.append(pair) + + +async def _probe_instance_connection( + conn_name: ConnectionName, + conn_info: ConnectionInfo, + enable_iam_auth: bool, + ip_type: IPTypes, + timeout: int, + iam_principals: list[tuple[str, str]], +) -> None: + """Proactively probes the database to refresh IAM tokens on server-side MCP.""" + targets: list[str] = [] + if conn_name.domain_name: + targets.append(conn_name.domain_name) + else: + if ip_type.value in conn_info.ip_addrs: + targets.extend(conn_info.ip_addrs[ip_type.value]) + + if not targets: + logger.debug( + f"['{conn_name}']: Proactive IAM token refresh probe skipped: no target IP addresses" + ) + return + + port = SERVER_PROXY_PORT + try: + ssl_context = await conn_info.create_ssl_context(enable_iam_auth) + except Exception as e: # noqa: BLE001 + logger.debug( + f"['{conn_name}']: Failed to create SSL context for probe: {e!s}" + ) + return + + principals: list[tuple[str, str] | None] = [] + if conn_info.database_version.startswith("POSTGRES") and iam_principals: + principals.extend(iam_principals) + else: + principals.append(None) + + all_succeeded = True + for principal in principals: + probed_principal = False + for target in targets: + try: + logger.debug( + f"['{conn_name}']: Probing IAM token refresh on {target}:{port}" + ) + reader, writer = await asyncio.wait_for( + asyncio.open_connection( + host=target, + port=port, + ssl=ssl_context, + server_hostname=( + conn_name.domain_name + if conn_name.domain_name + else None + ), + ), + timeout=float(timeout), + ) + try: + if principal is not None: + user, database = principal + writer.write( + _build_postgres_startup_packet(user, database) + ) + await writer.drain() + await asyncio.wait_for( + reader.read(1024), + timeout=float(timeout), + ) + writer.write(b"X\x00\x00\x00\x04") + await writer.drain() + finally: + writer.close() + await writer.wait_closed() + probed_principal = True + break + except Exception as e: # noqa: BLE001 + logger.debug( + f"['{conn_name}']: Probing IAM token refresh on {target}:{port} failed: {e!s}" + ) + if not probed_principal: + all_succeeded = False + + if all_succeeded: + logger.debug( + f"['{conn_name}']: Proactive IAM token refresh probe successful" + ) + else: + logger.debug( + f"['{conn_name}']: Proactive IAM token refresh probe encountered error across all targets" + ) + + +def _build_postgres_startup_packet(user: str, database: str) -> bytes: + """Builds a PostgreSQL v3.0 StartupMessage for (user, database).""" + if not database: + database = user + body = ( + b"user\x00" + + user.encode("utf-8") + + b"\x00database\x00" + + database.encode("utf-8") + + b"\x00\x00" + ) + total_len = 8 + len(body) + return total_len.to_bytes(4, byteorder="big") + b"\x00\x03\x00\x00" + body diff --git a/google/cloud/sql/connector/lazy.py b/google/cloud/sql/connector/lazy.py index 4e3de018..3a572ad4 100644 --- a/google/cloud/sql/connector/lazy.py +++ b/google/cloud/sql/connector/lazy.py @@ -24,6 +24,10 @@ from google.cloud.sql.connector.connection_info import ConnectionInfo from google.cloud.sql.connector.connection_info import ConnectionInfoCache from google.cloud.sql.connector.connection_name import ConnectionName +from google.cloud.sql.connector.enums import IPTypes +from google.cloud.sql.connector.instance import _append_iam_principal +from google.cloud.sql.connector.instance import _probe_instance_connection +from google.cloud.sql.connector.instance import DEFAULT_CONNECT_TIMEOUT from google.cloud.sql.connector.refresh_utils import _refresh_buffer logger = logging.getLogger(name=__name__) @@ -44,6 +48,8 @@ def __init__( client: CloudSQLClient, keys: asyncio.Future, enable_iam_auth: bool = False, + ip_type: IPTypes | str = IPTypes.PUBLIC, + timeout: int = DEFAULT_CONNECT_TIMEOUT, ) -> None: """Initializes a LazyRefreshCache instance. @@ -56,16 +62,27 @@ def __init__( enable_iam_auth (bool): Enables automatic IAM database authentication (Postgres and MySQL) as the default authentication method for all connections. + ip_type (IPTypes | str): Preferred IP type used to connect to the instance. + timeout (int): Connect timeout in seconds. """ self._conn_name = conn_name self._enable_iam_auth = enable_iam_auth + if isinstance(ip_type, str): + ip_type = IPTypes._from_str(ip_type) + self._ip_type = ip_type + self._timeout = timeout self._keys = keys self._client = client self._lock = asyncio.Lock() + self._iam_principals: list[tuple[str, str]] = [] self._cached: ConnectionInfo | None = None self._needs_refresh = False self._closed = False + def record_principal(self, user: str, database: str) -> None: + """Records a PostgreSQL (user, database) pair observed on an IAM connection.""" + _append_iam_principal(self._iam_principals, user, database) + @property def conn_name(self) -> ConnectionName: return self._conn_name @@ -82,6 +99,17 @@ async def force_refresh(self) -> None: async with self._lock: self._needs_refresh = True + async def _probe_connection(self, conn_info: ConnectionInfo) -> None: + """Proactively probes the database to refresh IAM tokens on server-side MCP.""" + await _probe_instance_connection( + self._conn_name, + conn_info, + self._enable_iam_auth, + self._ip_type, + self._timeout, + self._iam_principals, + ) + async def connect_info(self) -> ConnectionInfo: """Retrieves ConnectionInfo instance for establishing a secure connection to the Cloud SQL instance. @@ -110,6 +138,8 @@ async def connect_info(self) -> ConnectionInfo: self._keys, self._enable_iam_auth, ) + if self._enable_iam_auth: + await self._probe_connection(conn_info) except Exception as e: logger.debug( f"['{self._conn_name}']: Connection info " diff --git a/google/cloud/sql/connector/monitored_cache.py b/google/cloud/sql/connector/monitored_cache.py index d2ba2290..c88bced1 100644 --- a/google/cloud/sql/connector/monitored_cache.py +++ b/google/cloud/sql/connector/monitored_cache.py @@ -139,6 +139,9 @@ async def force_refresh(self) -> None: return return await self.cache.force_refresh() + def record_principal(self, user: str, database: str) -> None: + self.cache.record_principal(user, database) + async def close(self) -> None: # Cancel domain name ticker task. if self.domain_name_ticker: diff --git a/tests/unit/test_instance.py b/tests/unit/test_instance.py index ca87370b..9f5cfee1 100644 --- a/tests/unit/test_instance.py +++ b/tests/unit/test_instance.py @@ -391,3 +391,37 @@ async def test_ConnectionInfo_create_ssl_context_no_tls1_3_warning() -> None: in mock_logger.warning.call_args[0][0] ) + +@pytest.mark.asyncio +async def test_probe_connection_postgres_startup_packet( + cache: RefreshAheadCache, +) -> None: + """Test that _probe_connection sends PostgreSQL v3 StartupMessage and Terminate when IAM principal is recorded.""" + from unittest.mock import MagicMock + + from google.cloud.sql.connector.instance import _build_postgres_startup_packet + + cache._enable_iam_auth = True + cache.record_principal("iam-user@example.com", "mydb") + conn_info = await cache.connect_info() + + mock_reader = AsyncMock() + mock_reader.read = AsyncMock(return_value=b"R\x00\x00\x00\x08\x00\x00\x00\x00") + mock_writer = MagicMock() + mock_writer.drain = AsyncMock() + mock_writer.wait_closed = AsyncMock() + + with patch( + "google.cloud.sql.connector.instance.asyncio.open_connection", + AsyncMock(return_value=(mock_reader, mock_writer)), + ) as mock_open_conn: + await cache._probe_connection(conn_info) + + mock_open_conn.assert_awaited_once() + written_packets = [call.args[0] for call in mock_writer.write.call_args_list] + expected_startup = _build_postgres_startup_packet("iam-user@example.com", "mydb") + assert written_packets == [expected_startup, b"X\x00\x00\x00\x04"] + mock_reader.read.assert_awaited_once() + mock_writer.close.assert_called_once() + mock_writer.wait_closed.assert_awaited_once() + diff --git a/tests/unit/test_lazy.py b/tests/unit/test_lazy.py index 4b03ec05..63549a09 100644 --- a/tests/unit/test_lazy.py +++ b/tests/unit/test_lazy.py @@ -111,3 +111,47 @@ async def test_LazyRefreshCache_connect_info_error( await cache.close() + +async def test_LazyRefreshCache_probe_connection_postgres_startup_packet( + fake_client: CloudSQLClient, +) -> None: + """ + Test that LazyRefreshCache.connect_info probes the instance with a PostgreSQL + StartupMessage and Terminate when enable_iam_auth=True and a principal is recorded. + """ + from unittest.mock import MagicMock + from unittest.mock import patch + + from google.cloud.sql.connector.instance import _build_postgres_startup_packet + + keys = asyncio.create_task(generate_keys()) + cache = LazyRefreshCache( + ConnectionName("test-project", "test-region", "test-instance"), + client=fake_client, + keys=keys, + enable_iam_auth=True, + ) + cache.record_principal("iam-user@example.com", "mydb") + + mock_reader = AsyncMock() + mock_reader.read = AsyncMock(return_value=b"R\x00\x00\x00\x08\x00\x00\x00\x00") + mock_writer = MagicMock() + mock_writer.drain = AsyncMock() + mock_writer.wait_closed = AsyncMock() + + with patch( + "google.cloud.sql.connector.instance.asyncio.open_connection", + AsyncMock(return_value=(mock_reader, mock_writer)), + ) as mock_open_conn: + await cache.connect_info() + + mock_open_conn.assert_awaited_once() + written_packets = [call.args[0] for call in mock_writer.write.call_args_list] + expected_startup = _build_postgres_startup_packet("iam-user@example.com", "mydb") + assert written_packets == [expected_startup, b"X\x00\x00\x00\x04"] + mock_reader.read.assert_awaited_once() + mock_writer.close.assert_called_once() + mock_writer.wait_closed.assert_awaited_once() + await cache.close() + +