From bb86beb647818270f2c9948b28fa0176c42851b1 Mon Sep 17 00:00:00 2001 From: Jonathan Hess Date: Mon, 17 Aug 2026 23:52:52 +0000 Subject: [PATCH 1/5] feat: proactively probe database on Auto-IAM refresh to update MCP tokens --- google/cloud/sql/connector/instance.py | 59 ++++++++++++++++++++++++++ 1 file changed, 59 insertions(+) diff --git a/google/cloud/sql/connector/instance.py b/google/cloud/sql/connector/instance.py index 28ab54e4..a50c6c9d 100644 --- a/google/cloud/sql/connector/instance.py +++ b/google/cloud/sql/connector/instance.py @@ -119,6 +119,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 +140,63 @@ 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.""" + targets: list[str] = [] + if self._conn_name.domain_name: + targets.append(self._conn_name.domain_name) + else: + for ip_type in ("PSC", "PRIVATE", "PUBLIC"): + if ip_type in conn_info.ip_addrs: + targets.extend(conn_info.ip_addrs[ip_type]) + + if not targets: + logger.debug( + f"['{self._conn_name}']: Proactive IAM token refresh probe skipped: no target IP addresses" + ) + return + + port = 3307 + try: + ssl_context = await conn_info.create_ssl_context(self._enable_iam_auth) + except Exception as e: # noqa: BLE001 + logger.debug( + f"['{self._conn_name}']: Failed to create SSL context for probe: {e!s}" + ) + return + + for target in targets: + try: + logger.debug( + f"['{self._conn_name}']: Probing IAM token refresh on {target}:{port}" + ) + _, writer = await asyncio.wait_for( + asyncio.open_connection( + host=target, + port=port, + ssl=ssl_context, + server_hostname=( + self._conn_name.domain_name + if self._conn_name.domain_name + else None + ), + ), + timeout=15.0, + ) + writer.close() + await writer.wait_closed() + logger.debug( + f"['{self._conn_name}']: Proactive IAM token refresh probe successful" + ) + return + except Exception as e: # noqa: BLE001 + logger.debug( + f"['{self._conn_name}']: Probing IAM token refresh on {target}:{port} failed: {e!s}" + ) + logger.debug( + f"['{self._conn_name}']: Proactive IAM token refresh probe encountered error across all targets" + ) + def _schedule_refresh(self, delay: int) -> asyncio.Task: """ Schedule task to sleep and then perform refresh to get ConnectionInfo. From 46d1cf0ea6d15779a638a6d94a5843f91572bcaf Mon Sep 17 00:00:00 2001 From: Jonathan Hess Date: Tue, 18 Aug 2026 02:10:10 +0000 Subject: [PATCH 2/5] fix: address code review comments on proactive IAM token probe Code review comments addressed: - Respect the IP settings in connection configuration (self._ip_type) rather than probing all available IPs - Use SERVER_PROXY_PORT constant (3307) - Use DEFAULT_CONNECT_TIMEOUT constant (30) / configured timeout --- google/cloud/sql/connector/connector.py | 2 ++ google/cloud/sql/connector/instance.py | 20 +++++++++++++++----- 2 files changed, 17 insertions(+), 5 deletions(-) diff --git a/google/cloud/sql/connector/connector.py b/google/cloud/sql/connector/connector.py index 7a9964ca..ee362e5d 100644 --- a/google/cloud/sql/connector/connector.py +++ b/google/cloud/sql/connector/connector.py @@ -348,6 +348,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( diff --git a/google/cloud/sql/connector/instance.py b/google/cloud/sql/connector/instance.py index a50c6c9d..b226049f 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( @@ -146,9 +157,8 @@ async def _probe_connection(self, conn_info: ConnectionInfo) -> None: if self._conn_name.domain_name: targets.append(self._conn_name.domain_name) else: - for ip_type in ("PSC", "PRIVATE", "PUBLIC"): - if ip_type in conn_info.ip_addrs: - targets.extend(conn_info.ip_addrs[ip_type]) + if self._ip_type.value in conn_info.ip_addrs: + targets.extend(conn_info.ip_addrs[self._ip_type.value]) if not targets: logger.debug( @@ -156,7 +166,7 @@ async def _probe_connection(self, conn_info: ConnectionInfo) -> None: ) return - port = 3307 + port = SERVER_PROXY_PORT try: ssl_context = await conn_info.create_ssl_context(self._enable_iam_auth) except Exception as e: # noqa: BLE001 @@ -181,7 +191,7 @@ async def _probe_connection(self, conn_info: ConnectionInfo) -> None: else None ), ), - timeout=15.0, + timeout=float(self._timeout), ) writer.close() await writer.wait_closed() From 5a9b3f26d821f9672aa8b5a4be2882ae99ebfa29 Mon Sep 17 00:00:00 2001 From: kgala2 Date: Thu, 1 Oct 2026 06:50:51 +0000 Subject: [PATCH 3/5] fix: send PostgreSQL v3 StartupMessage in proactive IAM refresh probe --- google/cloud/sql/connector/connector.py | 3 + google/cloud/sql/connector/instance.py | 118 +++++++++++++----- google/cloud/sql/connector/monitored_cache.py | 4 + tests/unit/test_instance.py | 33 +++++ 4 files changed, 128 insertions(+), 30 deletions(-) diff --git a/google/cloud/sql/connector/connector.py b/google/cloud/sql/connector/connector.py index ee362e5d..6101c1f2 100644 --- a/google/cloud/sql/connector/connector.py +++ b/google/cloud/sql/connector/connector.py @@ -439,6 +439,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 b226049f..c1056483 100644 --- a/google/cloud/sql/connector/instance.py +++ b/google/cloud/sql/connector/instance.py @@ -84,10 +84,21 @@ 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.""" + if not user: + return + if not database: + database = user + pair = (user, database) + if pair not in self._iam_principals: + self._iam_principals.append(pair) + @property def conn_name(self) -> ConnectionName: return self._conn_name @@ -175,37 +186,69 @@ async def _probe_connection(self, conn_info: ConnectionInfo) -> None: ) return - for target in targets: - try: - logger.debug( - f"['{self._conn_name}']: Probing IAM token refresh on {target}:{port}" - ) - _, writer = await asyncio.wait_for( - asyncio.open_connection( - host=target, - port=port, - ssl=ssl_context, - server_hostname=( - self._conn_name.domain_name - if self._conn_name.domain_name - else None + principals: list[tuple[str, str] | None] = [] + if ( + conn_info.database_version.startswith("POSTGRES") + and self._iam_principals + ): + principals.extend(self._iam_principals) + else: + principals.append(None) + + all_succeeded = True + for principal in principals: + probed_principal = False + for target in targets: + try: + logger.debug( + f"['{self._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=( + self._conn_name.domain_name + if self._conn_name.domain_name + else None + ), ), - ), - timeout=float(self._timeout), - ) - writer.close() - await writer.wait_closed() - logger.debug( - f"['{self._conn_name}']: Proactive IAM token refresh probe successful" - ) - return - except Exception as e: # noqa: BLE001 - logger.debug( - f"['{self._conn_name}']: Probing IAM token refresh on {target}:{port} failed: {e!s}" - ) - logger.debug( - f"['{self._conn_name}']: Proactive IAM token refresh probe encountered error across all targets" - ) + timeout=float(self._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(self._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"['{self._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"['{self._conn_name}']: Proactive IAM token refresh probe successful" + ) + else: + logger.debug( + f"['{self._conn_name}']: Proactive IAM token refresh probe encountered error across all targets" + ) def _schedule_refresh(self, delay: int) -> asyncio.Task: """ @@ -292,3 +335,18 @@ 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 _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/monitored_cache.py b/google/cloud/sql/connector/monitored_cache.py index d2ba2290..f9d86f17 100644 --- a/google/cloud/sql/connector/monitored_cache.py +++ b/google/cloud/sql/connector/monitored_cache.py @@ -139,6 +139,10 @@ async def force_refresh(self) -> None: return return await self.cache.force_refresh() + def record_principal(self, user: str, database: str) -> None: + if hasattr(self.cache, "record_principal"): + 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..919747fc 100644 --- a/tests/unit/test_instance.py +++ b/tests/unit/test_instance.py @@ -391,3 +391,36 @@ 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() + From 8a44dc73f8ebdb3b2385dcf0d49f55eb6b62f389 Mon Sep 17 00:00:00 2001 From: kgala2 Date: Thu, 1 Oct 2026 21:04:39 +0000 Subject: [PATCH 4/5] refactor: add record_principal to ConnectionInfoCache ABC and probe in LazyRefreshCache --- google/cloud/sql/connector/connection_info.py | 4 + google/cloud/sql/connector/connector.py | 2 + google/cloud/sql/connector/instance.py | 207 ++++++++++-------- google/cloud/sql/connector/lazy.py | 30 +++ google/cloud/sql/connector/monitored_cache.py | 3 +- tests/unit/test_lazy.py | 42 ++++ 6 files changed, 194 insertions(+), 94 deletions(-) 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 6101c1f2..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( diff --git a/google/cloud/sql/connector/instance.py b/google/cloud/sql/connector/instance.py index c1056483..eeb297b2 100644 --- a/google/cloud/sql/connector/instance.py +++ b/google/cloud/sql/connector/instance.py @@ -91,13 +91,7 @@ def __init__( def record_principal(self, user: str, database: str) -> None: """Records a PostgreSQL (user, database) pair observed on an IAM connection.""" - if not user: - return - if not database: - database = user - pair = (user, database) - if pair not in self._iam_principals: - self._iam_principals.append(pair) + _append_iam_principal(self._iam_principals, user, database) @property def conn_name(self) -> ConnectionName: @@ -164,91 +158,14 @@ async def _perform_refresh(self) -> ConnectionInfo: async def _probe_connection(self, conn_info: ConnectionInfo) -> None: """Proactively probes the database to refresh IAM tokens on server-side MCP.""" - targets: list[str] = [] - if self._conn_name.domain_name: - targets.append(self._conn_name.domain_name) - else: - if self._ip_type.value in conn_info.ip_addrs: - targets.extend(conn_info.ip_addrs[self._ip_type.value]) - - if not targets: - logger.debug( - f"['{self._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(self._enable_iam_auth) - except Exception as e: # noqa: BLE001 - logger.debug( - f"['{self._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 self._iam_principals - ): - principals.extend(self._iam_principals) - else: - principals.append(None) - - all_succeeded = True - for principal in principals: - probed_principal = False - for target in targets: - try: - logger.debug( - f"['{self._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=( - self._conn_name.domain_name - if self._conn_name.domain_name - else None - ), - ), - timeout=float(self._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(self._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"['{self._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"['{self._conn_name}']: Proactive IAM token refresh probe successful" - ) - else: - logger.debug( - f"['{self._conn_name}']: Proactive IAM token refresh probe encountered error across all targets" - ) + 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: """ @@ -337,6 +254,112 @@ async def close(self) -> None: 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: 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 f9d86f17..c88bced1 100644 --- a/google/cloud/sql/connector/monitored_cache.py +++ b/google/cloud/sql/connector/monitored_cache.py @@ -140,8 +140,7 @@ async def force_refresh(self) -> None: return await self.cache.force_refresh() def record_principal(self, user: str, database: str) -> None: - if hasattr(self.cache, "record_principal"): - self.cache.record_principal(user, database) + self.cache.record_principal(user, database) async def close(self) -> None: # Cancel domain name ticker task. diff --git a/tests/unit/test_lazy.py b/tests/unit/test_lazy.py index 4b03ec05..c86f91eb 100644 --- a/tests/unit/test_lazy.py +++ b/tests/unit/test_lazy.py @@ -111,3 +111,45 @@ 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, 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() + + From 8f22f4d34462e06c7f6075d137c98991b296c32b Mon Sep 17 00:00:00 2001 From: kgala2 Date: Thu, 1 Oct 2026 23:20:12 +0000 Subject: [PATCH 5/5] style: sort test imports with ruff --- tests/unit/test_instance.py | 1 + tests/unit/test_lazy.py | 4 +++- 2 files changed, 4 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_instance.py b/tests/unit/test_instance.py index 919747fc..9f5cfee1 100644 --- a/tests/unit/test_instance.py +++ b/tests/unit/test_instance.py @@ -398,6 +398,7 @@ async def test_probe_connection_postgres_startup_packet( ) -> 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 diff --git a/tests/unit/test_lazy.py b/tests/unit/test_lazy.py index c86f91eb..63549a09 100644 --- a/tests/unit/test_lazy.py +++ b/tests/unit/test_lazy.py @@ -119,7 +119,9 @@ async def test_LazyRefreshCache_probe_connection_postgres_startup_packet( 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, patch + 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())