Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions google/cloud/sql/connector/connection_info.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
7 changes: 7 additions & 0 deletions google/cloud/sql/connector/connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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(
Expand Down Expand Up @@ -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
Expand Down
150 changes: 150 additions & 0 deletions google/cloud/sql/connector/instance.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):
Expand All @@ -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.

Expand All @@ -62,21 +67,32 @@ 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(
max_capacity=2,
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
Expand Down Expand Up @@ -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"
)
Expand All @@ -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.
Expand Down Expand Up @@ -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
30 changes: 30 additions & 0 deletions google/cloud/sql/connector/lazy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand All @@ -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.

Expand All @@ -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
Expand All @@ -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.
Expand Down Expand Up @@ -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 "
Expand Down
3 changes: 3 additions & 0 deletions google/cloud/sql/connector/monitored_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
34 changes: 34 additions & 0 deletions tests/unit/test_instance.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Loading
Loading