diff --git a/.coveragerc b/.coveragerc new file mode 100644 index 00000000..bbdd152b --- /dev/null +++ b/.coveragerc @@ -0,0 +1,10 @@ +[report] +exclude_lines = + pragma: no cover + def __repr__ + raise AssertionError + raise NotImplementedError + if __name__ == .__main__.: + @abc.abstractmethod + if TYPE_CHECKING: + diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index feb74526..2cfbd004 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -86,6 +86,14 @@ jobs: POSTGRES_CUSTOMER_CAS_PASS_VALID_DOMAIN_NAME:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_CUSTOMER_CAS_PASS_VALID_DOMAIN_NAME POSTGRES_MCP_CONNECTION_NAME:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_MCP_CONNECTION_NAME POSTGRES_MCP_PASS:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_MCP_PASS + POSTGRES_AIDE_CONNECTION_NAME:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_AIDE_CONNECTION_NAME + POSTGRES_AIDE_USER:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_AIDE_USER + POSTGRES_AIDE_PASS:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_AIDE_PASS + POSTGRES_AIDE_DB:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_AIDE_DB + POSTGRES_FALLBACK_CONNECTION_NAME:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_FALLBACK_CONNECTION_NAME + POSTGRES_FALLBACK_USER:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_FALLBACK_USER + POSTGRES_FALLBACK_PASS:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_FALLBACK_PASS + POSTGRES_FALLBACK_DB:${{ vars.GOOGLE_CLOUD_PROJECT }}/POSTGRES_FALLBACK_DB SQLSERVER_CONNECTION_NAME:${{ vars.GOOGLE_CLOUD_PROJECT }}/SQLSERVER_CONNECTION_NAME SQLSERVER_USER:${{ vars.GOOGLE_CLOUD_PROJECT }}/SQLSERVER_USER SQLSERVER_PASS:${{ vars.GOOGLE_CLOUD_PROJECT }}/SQLSERVER_PASS @@ -112,6 +120,14 @@ jobs: POSTGRES_CUSTOMER_CAS_PASS_VALID_DOMAIN_NAME: "${{ steps.secrets.outputs.POSTGRES_CUSTOMER_CAS_PASS_VALID_DOMAIN_NAME }}" POSTGRES_MCP_CONNECTION_NAME: "${{ steps.secrets.outputs.POSTGRES_MCP_CONNECTION_NAME }}" POSTGRES_MCP_PASS: "${{ steps.secrets.outputs.POSTGRES_MCP_PASS }}" + POSTGRES_AIDE_CONNECTION_NAME: "${{ steps.secrets.outputs.POSTGRES_AIDE_CONNECTION_NAME }}" + POSTGRES_AIDE_USER: "${{ steps.secrets.outputs.POSTGRES_AIDE_USER }}" + POSTGRES_AIDE_PASS: "${{ steps.secrets.outputs.POSTGRES_AIDE_PASS }}" + POSTGRES_AIDE_DB: "${{ steps.secrets.outputs.POSTGRES_AIDE_DB }}" + POSTGRES_FALLBACK_CONNECTION_NAME: "${{ steps.secrets.outputs.POSTGRES_FALLBACK_CONNECTION_NAME }}" + POSTGRES_FALLBACK_USER: "${{ steps.secrets.outputs.POSTGRES_FALLBACK_USER }}" + POSTGRES_FALLBACK_PASS: "${{ steps.secrets.outputs.POSTGRES_FALLBACK_PASS }}" + POSTGRES_FALLBACK_DB: "${{ steps.secrets.outputs.POSTGRES_FALLBACK_DB }}" SQLSERVER_CONNECTION_NAME: "${{ steps.secrets.outputs.SQLSERVER_CONNECTION_NAME }}" SQLSERVER_USER: "${{ steps.secrets.outputs.SQLSERVER_USER }}" SQLSERVER_PASS: "${{ steps.secrets.outputs.SQLSERVER_PASS }}" diff --git a/.gitignore b/.gitignore index 6fed8e21..a8ce0684 100644 --- a/.gitignore +++ b/.gitignore @@ -10,6 +10,7 @@ dist/ sponge_log.xml .envrc *.iml +build/ .mypy_cache/ .nox/ .pytest_cache/ diff --git a/build.sh b/build.sh index 3c76d992..2987465e 100755 --- a/build.sh +++ b/build.sh @@ -120,6 +120,14 @@ function write_e2e_env(){ POSTGRES_CUSTOMER_CAS_INVALID_DOMAIN_NAME=POSTGRES_CUSTOMER_CAS_INVALID_DOMAIN_NAME POSTGRES_MCP_CONNECTION_NAME=POSTGRES_MCP_CONNECTION_NAME POSTGRES_MCP_PASS=POSTGRES_MCP_PASS + POSTGRES_AIDE_CONNECTION_NAME=POSTGRES_AIDE_CONNECTION_NAME + POSTGRES_AIDE_USER=POSTGRES_AIDE_USER + POSTGRES_AIDE_PASS=POSTGRES_AIDE_PASS + POSTGRES_AIDE_DB=POSTGRES_AIDE_DB + POSTGRES_FALLBACK_CONNECTION_NAME=POSTGRES_FALLBACK_CONNECTION_NAME + POSTGRES_FALLBACK_USER=POSTGRES_FALLBACK_USER + POSTGRES_FALLBACK_PASS=POSTGRES_FALLBACK_PASS + POSTGRES_FALLBACK_DB=POSTGRES_FALLBACK_DB SQLSERVER_CONNECTION_NAME=SQLSERVER_CONNECTION_NAME SQLSERVER_USER=SQLSERVER_USER SQLSERVER_PASS=SQLSERVER_PASS diff --git a/google/cloud/sql/connector/asyncpg.py b/google/cloud/sql/connector/asyncpg.py index 191a5618..07c8e98f 100644 --- a/google/cloud/sql/connector/asyncpg.py +++ b/google/cloud/sql/connector/asyncpg.py @@ -14,6 +14,8 @@ limitations under the License. """ +from __future__ import annotations + import ssl from typing import Any, TYPE_CHECKING @@ -24,16 +26,15 @@ async def connect( - ip_address: str, ctx: ssl.SSLContext, **kwargs: Any -) -> "asyncpg.Connection": + ip_address: str, ctx: ssl.SSLContext | None, **kwargs: Any +) -> asyncpg.Connection: """Helper function to create an asyncpg DB-API connection object. Args: ip_address (str): A string containing an IP address for the Cloud SQL instance. ctx (ssl.SSLContext): An SSLContext object created from the Cloud SQL - server CA cert and ephemeral cert. - server CA cert and ephemeral cert. + server CA cert and ephemeral cert. Pass None to disable SSL. kwargs: Keyword arguments for establishing asyncpg connection object to Cloud SQL instance. @@ -55,14 +56,18 @@ async def connect( if db is None: raise KeyError("database") passwd = kwargs.pop("password", None) + port = kwargs.pop("port", SERVER_PROXY_PORT) - return await asyncpg.connect( - user=user, - database=db, - password=passwd, - host=ip_address, - port=SERVER_PROXY_PORT, - ssl=ctx, - direct_tls=True, + connect_args = { + "user": user, + "database": db, + "password": passwd, + "host": ip_address, + "port": port, **kwargs, - ) + } + if ctx is not None: + connect_args["ssl"] = ctx + connect_args["direct_tls"] = True + + return await asyncpg.connect(**connect_args) diff --git a/google/cloud/sql/connector/client.py b/google/cloud/sql/connector/client.py index ebf823e2..b0f83b6d 100644 --- a/google/cloud/sql/connector/client.py +++ b/google/cloud/sql/connector/client.py @@ -173,9 +173,13 @@ async def _get_metadata( if psc_dns_names: ip_addresses["PSC"] = psc_dns_names + server_ca_cert = None + if "serverCaCert" in ret_dict and "cert" in ret_dict["serverCaCert"]: + server_ca_cert = ret_dict["serverCaCert"]["cert"] + return { "ip_addresses": ip_addresses, - "server_ca_cert": ret_dict["serverCaCert"]["cert"], + "server_ca_cert": server_ca_cert, "database_version": ret_dict["databaseVersion"], } @@ -271,7 +275,15 @@ async def _get_ephemeral( finally: resp.raise_for_status() - ephemeral_cert: str = ret_dict["ephemeralCert"]["cert"] + try: + ephemeral_cert: str = ret_dict["ephemeralCert"]["cert"] + except KeyError as e: + logger.error( + "KeyError in _get_ephemeral parsing generateEphemeralCert: %s. Response dict: %s", + e, + ret_dict, + ) + raise # decode cert to read expiration x509 = load_pem_x509_certificate( diff --git a/google/cloud/sql/connector/connection_info.py b/google/cloud/sql/connector/connection_info.py index 49404f49..f0b8f1dd 100644 --- a/google/cloud/sql/connector/connection_info.py +++ b/google/cloud/sql/connector/connection_info.py @@ -21,6 +21,7 @@ from typing import Any, TYPE_CHECKING from google.cloud.sql.connector.connection_name import ConnectionName +from google.cloud.sql.connector.exceptions import CloudSQLConnectionError from google.cloud.sql.connector.exceptions import CloudSQLIPTypeError from google.cloud.sql.connector.exceptions import TLSVersionError from google.cloud.sql.connector.utils import AsyncTemporaryDirectory @@ -62,7 +63,7 @@ class ConnectionInfo: conn_name: ConnectionName client_cert: str - server_ca_cert: str + server_ca_cert: str | None private_key: bytes ip_addrs: dict[str, Any] database_version: str @@ -78,6 +79,12 @@ async def create_ssl_context(self, enable_iam_auth: bool = False) -> ssl.SSLCont # if SSL context is cached, use it if self.context is not None: return self.context + + if self.server_ca_cert is None: + raise CloudSQLConnectionError( + "Cannot create SSL context: server CA certificate is missing." + ) + context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) # update ssl.PROTOCOL_TLS_CLIENT default diff --git a/google/cloud/sql/connector/connector.py b/google/cloud/sql/connector/connector.py index 7a9964ca..5f89183b 100644 --- a/google/cloud/sql/connector/connector.py +++ b/google/cloud/sql/connector/connector.py @@ -20,8 +20,10 @@ from functools import partial import logging import os +import random import socket from threading import Thread +import time from types import TracebackType from typing import Any, Callable @@ -35,16 +37,20 @@ from google.cloud.sql.connector import pymysql from google.cloud.sql.connector import pytds from google.cloud.sql.connector.client import CloudSQLClient +from google.cloud.sql.connector.connection_name import ConnectionName from google.cloud.sql.connector.enums import DriverMapping from google.cloud.sql.connector.enums import IPTypes from google.cloud.sql.connector.enums import RefreshStrategy from google.cloud.sql.connector.exceptions import ClosedConnectorError from google.cloud.sql.connector.exceptions import ConnectorLoopError +from google.cloud.sql.connector.exceptions import IncompatibleDriverError +from google.cloud.sql.connector.exceptions import ResourceExhaustedError from google.cloud.sql.connector.instance import RefreshAheadCache from google.cloud.sql.connector.lazy import LazyRefreshCache from google.cloud.sql.connector.monitored_cache import MonitoredCache from google.cloud.sql.connector.resolver import DefaultResolver from google.cloud.sql.connector.resolver import DnsResolver +from google.cloud.sql.connector.sqldata_client import SqlDataClient from google.cloud.sql.connector.utils import format_database_user from google.cloud.sql.connector.utils import generate_keys @@ -57,6 +63,49 @@ _SQLADMIN_HOST_TEMPLATE = "sqladmin.{universe_domain}" +class SqlDataConnState: + """Tracks connection state, fallback status, and resource exhaustion cooldown for SQL Data Service.""" + + def __init__(self) -> None: + self.allowed: bool = True + self.cooldown_until: float | None = None + self.backoff_counter: int = 0 + self.last_err: Exception | None = None + + def is_cooldown_active(self) -> bool: + """Returns True if the instance connection is in cooldown due to resource exhaustion.""" + return bool( + self.allowed + and self.cooldown_until + and time.time() < self.cooldown_until + ) + + def record_exhausted(self, err: Exception, base_cooldown: float) -> float: + """Records a resource exhaustion error, increments backoff counter, and returns the cooldown duration.""" + if self.backoff_counter < 5: + self.backoff_counter += 1 + backoff = _cooldown_backoff(base_cooldown, self.backoff_counter) + self.cooldown_until = time.time() + backoff + self.last_err = err + return backoff + + def record_success(self) -> None: + """Resets cooldown and backoff state on successful communication.""" + self.backoff_counter = 0 + self.cooldown_until = None + self.last_err = None + + def record_fallback(self) -> None: + """Marks SQL Data Service as not allowed for this instance.""" + self.allowed = False + + +def _cooldown_backoff(base_cooldown: float, attempt: int) -> float: + multi = 1.618 + exp = float(attempt - 1) + random.random() + return base_cooldown * (multi**exp) + + class Connector: """Configure and create secure connections to Cloud SQL.""" @@ -74,14 +123,18 @@ def __init__( refresh_strategy: str | RefreshStrategy = RefreshStrategy.BACKGROUND, resolver: type[DefaultResolver | DnsResolver] = DefaultResolver, failover_period: int = 30, + sql_data_endpoint: str = "sqladmin.googleapis.com", + sql_data_stream_timeout: int = 7200, + resource_exhausted_cooldown_period: float = 5.0, ) -> None: """Initializes a Connector instance. Args: ip_type (str | IPTypes): The default IP address type used to connect to Cloud SQL instances. Can be one of the following: - IPTypes.PUBLIC ("PUBLIC"), IPTypes.PRIVATE ("PRIVATE"), or - IPTypes.PSC ("PSC"). Default: IPTypes.PUBLIC + IPTypes.PUBLIC ("PUBLIC"), IPTypes.PRIVATE ("PRIVATE"), + IPTypes.PSC ("PSC"), or IPTypes.SQL_DATA ("SQL_DATA"). + Default: IPTypes.PUBLIC enable_iam_auth (bool): Enables automatic IAM database authentication (Postgres and MySQL) as the default authentication method for all @@ -126,6 +179,15 @@ def __init__( attempt to check if a failover has occured for a given instance. Must be used with `resolver=DnsResolver` to have any effect. Default: 30 + + sql_data_endpoint (str): Endpoint host for SQL Data Service calls. + Default: "sqladmin.googleapis.com". + + sql_data_stream_timeout (int): Timeout in seconds for the SQL Data + Service gRPC stream. Default: 7200. + + resource_exhausted_cooldown_period (float): Cooldown period in seconds + after a ResourceExhausted error. Default: 5.0. """ # if refresh_strategy is str, convert to RefreshStrategy enum if isinstance(refresh_strategy, str): @@ -214,6 +276,16 @@ def __init__( "configured the universe domain explicitly, `googleapis.com` " "is the default." ) + self._sql_data_endpoint = sql_data_endpoint + self._sql_data_stream_timeout = sql_data_stream_timeout + self._resource_exhausted_cooldown_period = ( + resource_exhausted_cooldown_period + ) + self._sql_data_fallback_cache: set[str] = set() + self._sql_data_conn_state: dict[str, SqlDataConnState] = {} + self._sqldata_clients: set[Any] = set() + + @property def universe_domain(self) -> str: @@ -260,6 +332,49 @@ def connect( ) return connect_future.result() + def _get_or_create_cache( + self, + conn_name: ConnectionName, + enable_iam_auth: bool, + ) -> MonitoredCache: + assert self._client is not None, "client must be initialized before creating cache" + assert self._keys is not None, "keys must be initialized before creating cache" + assert self._resolver is not None, "resolver must be initialized before creating cache" + if (str(conn_name), enable_iam_auth) in self._cache and not self._cache[ + (str(conn_name), enable_iam_auth) + ].closed: + return self._cache[(str(conn_name), enable_iam_auth)] + + if self._refresh_strategy == RefreshStrategy.LAZY: + logger.debug( + f"['{conn_name}']: Refresh strategy is set to lazy refresh" + ) + cache: LazyRefreshCache | RefreshAheadCache = LazyRefreshCache( + conn_name, + self._client, + self._keys, + enable_iam_auth, + ) + else: + logger.debug( + f"['{conn_name}']: Refresh strategy is set to backgound refresh" + ) + cache = RefreshAheadCache( + conn_name, + self._client, + self._keys, + enable_iam_auth, + ) + # wrap cache as a MonitoredCache + monitored_cache = MonitoredCache( + cache, + self._failover_period, + self._resolver, + ) + logger.debug(f"['{conn_name}']: Connection info added to cache") + self._cache[(str(conn_name), enable_iam_auth)] = monitored_cache + return monitored_cache + async def connect_async( self, instance_connection_string: str, driver: str, **kwargs: Any ) -> Any: @@ -321,42 +436,13 @@ async def connect_async( if self._resolver is None: self._resolver = self._resolver_cls(client=self._client) enable_iam_auth = kwargs.pop("enable_iam_auth", self._enable_iam_auth) + ip_type = kwargs.pop("ip_type", self._ip_type) + if isinstance(ip_type, str): + ip_type = IPTypes._from_str(ip_type) conn_name = await self._resolver.resolve(instance_connection_string) - # Cache entry must exist and not be closed - if (str(conn_name), enable_iam_auth) in self._cache and not self._cache[ - (str(conn_name), enable_iam_auth) - ].closed: - monitored_cache = self._cache[(str(conn_name), enable_iam_auth)] - else: - if self._refresh_strategy == RefreshStrategy.LAZY: - logger.debug( - f"['{conn_name}']: Refresh strategy is set to lazy refresh" - ) - cache: LazyRefreshCache | RefreshAheadCache = LazyRefreshCache( - conn_name, - self._client, - self._keys, - enable_iam_auth, - ) - else: - logger.debug( - f"['{conn_name}']: Refresh strategy is set to backgound refresh" - ) - cache = RefreshAheadCache( - conn_name, - self._client, - self._keys, - enable_iam_auth, - ) - # wrap cache as a MonitoredCache - monitored_cache = MonitoredCache( - cache, - self._failover_period, - self._resolver, - ) - logger.debug(f"['{conn_name}']: Connection info added to cache") - self._cache[(str(conn_name), enable_iam_auth)] = monitored_cache + if ip_type != IPTypes.SQL_DATA: + monitored_cache = self._get_or_create_cache(conn_name, enable_iam_auth) connect_func = { "pymysql": pymysql.connect, @@ -371,11 +457,6 @@ async def connect_async( connector: Callable = connect_func[driver] # type: ignore except KeyError: raise KeyError(f"Driver '{driver}' is not supported.") - - ip_type = kwargs.pop("ip_type", self._ip_type) - # if ip_type is str, convert to IPTypes enum - if isinstance(ip_type, str): - ip_type = IPTypes._from_str(ip_type) kwargs["timeout"] = kwargs.get("timeout", self._timeout) # Host and ssl options come from the certificates and metadata, so we don't @@ -384,113 +465,207 @@ async def connect_async( kwargs.pop("ssl", None) kwargs.pop("port", None) - # attempt to get connection info for Cloud SQL instance + # attempt to establish connection try: - conn_info = await monitored_cache.connect_info() - # validate driver matches intended database engine - DriverMapping.validate_engine(driver, conn_info.database_version) - preferred_ips = conn_info.get_preferred_ips(ip_type) - except Exception: - # with an error from Cloud SQL Admin API call or IP type, invalidate - # the cache and re-raise the error - await self._remove_cached(str(conn_name), enable_iam_auth) - raise - - targets = [] - # If the connector is configured with a custom DNS name, attempt to use - # that DNS name to connect to the instance. Fall back to the metadata IP - # address if the DNS name does not resolve to an IP address. - if conn_info.conn_name.domain_name and isinstance(self._resolver, DnsResolver): - try: - ips = await self._resolver.resolve_a_record(conn_info.conn_name.domain_name) - if ips: - targets.extend(ips) - logger.debug( - f"['{instance_connection_string}']: Custom DNS name " - f"'{conn_info.conn_name.domain_name}' resolved to '{ips}', " - "using it to connect" + if ip_type == IPTypes.SQL_DATA: + if driver in ASYNC_DRIVERS: + raise IncompatibleDriverError( + f"Driver '{driver}' is not supported with ip_type '{ip_type}'." ) - else: + state = self._sql_data_conn_state.setdefault( + str(conn_name), SqlDataConnState() + ) + if state.is_cooldown_active(): logger.debug( - f"['{instance_connection_string}']: Custom DNS name " - f"'{conn_info.conn_name.domain_name}' resolved but returned no " - f"entries, using '{preferred_ips}' from instance metadata" + f"['{conn_name}']: SQL Data Service in cooldown until {state.cooldown_until}" ) - targets.extend(preferred_ips) - except Exception as e: # noqa: BLE001 - logger.debug( - f"['{instance_connection_string}']: Custom DNS name " - f"'{conn_info.conn_name.domain_name}' did not resolve to an IP " - f"address: {e}, using '{preferred_ips}' from instance metadata" + raise ResourceExhaustedError( + "cooldown active", str(conn_name), state.last_err + ) + + logger.debug(f"['{conn_name}']: Connecting via SQL Data Service tunnel") + if enable_iam_auth: + engine = DriverMapping[driver.upper()].value + formatted_user = format_database_user( + engine, kwargs["user"] + ) + if formatted_user != kwargs["user"]: + logger.debug( + f"['{instance_connection_string}']: Truncated IAM database username from {kwargs['user']} to {formatted_user}" + ) + kwargs["user"] = formatted_user + + sqldata_client = SqlDataClient( + endpoint=self._sql_data_endpoint, + credentials=self._credentials, + quota_project=self._quota_project, + timeout=self._sql_data_stream_timeout, + ) + self._sqldata_clients.add(sqldata_client) + sqldata_client._on_close_callbacks.append( + lambda: self._sqldata_clients.discard(sqldata_client) ) - targets.extend(preferred_ips) - else: - targets.extend(preferred_ips) - # format `user` param for automatic IAM database authn - if enable_iam_auth: - formatted_user = format_database_user( - conn_info.database_version, kwargs["user"] - ) - if formatted_user != kwargs["user"]: - logger.debug( - f"['{instance_connection_string}']: Truncated IAM database username from {kwargs['user']} to {formatted_user}" + def on_resource_exhausted(err: Exception) -> None: + backoff = state.record_exhausted( + err, self._resource_exhausted_cooldown_period + ) + logger.debug( + f"['{conn_name}']: ResourceExhausted occurred, backing off for {backoff:.2f}s " + f"(attempt {state.backoff_counter})" + ) + + def on_success() -> None: + state.record_success() + + def on_fallback(name: str) -> None: + state.record_fallback() + self._sql_data_fallback_cache.add(name) + + def is_fallback_cached(name: str) -> bool: + return not state.allowed or name in self._sql_data_fallback_cache + + # Defer cache creation and connect_info call + async def get_conn_info(): + cache = self._get_or_create_cache(conn_name, enable_iam_auth) + return await cache.connect_info() + + sock = await sqldata_client.connect( + instance_connection_name=str(conn_name), + region=conn_name.region, + project=conn_name.project, + get_conn_info=get_conn_info, + enable_iam_auth=enable_iam_auth, + on_fallback=on_fallback, + is_fallback_cached=is_fallback_cached, + on_resource_exhausted=on_resource_exhausted, + on_success=on_success, + connect_timeout=kwargs.get("timeout", self._timeout), ) - kwargs["user"] = formatted_user - try: - last_ex = None - for target_ip in targets: - logger.debug(f"['{conn_info.conn_name}']: Connecting to {target_ip}:3307") + cache = None + if conn_name.domain_name: + cache = self._get_or_create_cache(conn_name, enable_iam_auth) + cache.sockets.append(sock) + + connect_partial = partial( + connector, + "127.0.0.1", + sock, + **kwargs, + ) + try: + return await self._loop.run_in_executor(None, connect_partial) + except Exception: + sock.close() + if conn_name.domain_name and cache: + cache._purge_closed_sockets() + raise + else: try: - # async drivers are unblocking and can be awaited directly - if driver in ASYNC_DRIVERS: - conn = await connector( + conn_info = await monitored_cache.connect_info() + # validate driver matches intended database engine + DriverMapping.validate_engine(driver, conn_info.database_version) + preferred_ips = conn_info.get_preferred_ips(ip_type) + except Exception: + # with an error from Cloud SQL Admin API call or IP type, invalidate + # the cache and re-raise the error + await self._remove_cached(str(conn_name), enable_iam_auth) + raise + + targets = [] + if conn_info.conn_name.domain_name and isinstance(self._resolver, DnsResolver): + try: + ips = await self._resolver.resolve_a_record(conn_info.conn_name.domain_name) + if ips: + targets.extend(ips) + logger.debug( + f"['{instance_connection_string}']: Custom DNS name " + f"'{conn_info.conn_name.domain_name}' resolved to '{ips}', " + "using it to connect" + ) + else: + logger.debug( + f"['{instance_connection_string}']: Custom DNS name " + f"'{conn_info.conn_name.domain_name}' resolved but returned no " + f"entries, using '{preferred_ips}' from instance metadata" + ) + targets.extend(preferred_ips) + except Exception as e: # noqa: BLE001 + logger.debug( + f"['{instance_connection_string}']: Custom DNS name " + f"'{conn_info.conn_name.domain_name}' did not resolve to an IP " + f"address: {e}, using '{preferred_ips}' from instance metadata" + ) + targets.extend(preferred_ips) + else: + targets.extend(preferred_ips) + + # format `user` param for automatic IAM database authn + if enable_iam_auth: + formatted_user = format_database_user( + conn_info.database_version, kwargs["user"] + ) + if formatted_user != kwargs["user"]: + logger.debug( + f"['{instance_connection_string}']: Truncated IAM database username from {kwargs['user']} to {formatted_user}" + ) + kwargs["user"] = formatted_user + + last_ex = None + for target_ip in targets: + logger.debug(f"['{conn_info.conn_name}']: Connecting to {target_ip}:3307") + try: + # async drivers are unblocking and can be awaited directly + if driver in ASYNC_DRIVERS: + conn = await connector( + target_ip, + await conn_info.create_ssl_context(enable_iam_auth), + **kwargs, + ) + last_ex = None + return conn + + # Create socket with SSLContext for sync drivers + ctx = await conn_info.create_ssl_context(enable_iam_auth) + raw_sock = socket.create_connection((target_ip, SERVER_PROXY_PORT)) + try: + sock = ctx.wrap_socket( + raw_sock, + server_hostname=target_ip, + ) + except Exception: + raw_sock.close() + raise + + # If this connection was opened using a domain name, then store it + # for later in case we need to forcibly close it on failover. + if conn_info.conn_name.domain_name: + monitored_cache.sockets.append(sock) + # Synchronous drivers are blocking and run using executor + connect_partial = partial( + connector, target_ip, - await conn_info.create_ssl_context(enable_iam_auth), + sock, **kwargs, ) + conn = await self._loop.run_in_executor(None, connect_partial) last_ex = None return conn - - # Create socket with SSLContext for sync drivers - ctx = await conn_info.create_ssl_context(enable_iam_auth) - raw_sock = socket.create_connection((target_ip, SERVER_PROXY_PORT)) - try: - sock = ctx.wrap_socket( - raw_sock, - server_hostname=target_ip, + except Exception as e: # noqa: BLE001 + logger.debug( + f"['{conn_info.conn_name}']: Connection to {target_ip} failed: {e}" ) - except Exception: - raw_sock.close() - raise - - # If this connection was opened using a domain name, then store it - # for later in case we need to forcibly close it on failover. - if conn_info.conn_name.domain_name: - monitored_cache.sockets.append(sock) - # Synchronous drivers are blocking and run using executor - connect_partial = partial( - connector, - target_ip, - sock, - **kwargs, - ) - conn = await self._loop.run_in_executor(None, connect_partial) - last_ex = None - return conn - except Exception as e: # noqa: BLE001 - logger.debug( - f"['{conn_info.conn_name}']: Connection to {target_ip} failed: {e}" - ) - last_ex = e + last_ex = e - if last_ex: - raise last_ex + if last_ex: + raise last_ex except Exception: # with any exception, we attempt a force refresh, then throw the error - await monitored_cache.force_refresh() + cached_entry = self._cache.get((str(conn_name), enable_iam_auth)) + if cached_entry: + await cached_entry.force_refresh() raise async def _remove_cached( @@ -538,8 +713,11 @@ def close(self) -> None: close_future = asyncio.run_coroutine_threadsafe( self.close_async(), loop=self._loop ) - # Will attempt to safely shut down tasks for 3s - close_future.result(timeout=3) + try: + # Will attempt to safely shut down tasks for 3s + close_future.result(timeout=3) + except Exception as e: # noqa: BLE001 + logger.error(f"Error during close_async: {e}") # if background thread exists for Connector, clean it up if self._thread: if self._loop.is_running(): @@ -554,7 +732,11 @@ async def close_async(self) -> None: self._closed = True if self._client: await self._client.close() - await asyncio.gather(*[cache.close() for cache in self._cache.values()]) + await asyncio.gather( + *[cache.close() for cache in self._cache.values()], + *[client.close() for client in list(self._sqldata_clients)], + return_exceptions=True, + ) async def create_async_connector( @@ -570,6 +752,9 @@ async def create_async_connector( refresh_strategy: str | RefreshStrategy = RefreshStrategy.BACKGROUND, resolver: type[DefaultResolver | DnsResolver] = DefaultResolver, failover_period: int = 30, + sql_data_endpoint: str = "sqladmin.googleapis.com", + sql_data_stream_timeout: int = 7200, + resource_exhausted_cooldown_period: float = 5.0, ) -> Connector: """Helper function to create Connector object for asyncio connections. @@ -579,8 +764,9 @@ async def create_async_connector( Args: ip_type (str | IPTypes): The default IP address type used to connect to Cloud SQL instances. Can be one of the following: - IPTypes.PUBLIC ("PUBLIC"), IPTypes.PRIVATE ("PRIVATE"), or - IPTypes.PSC ("PSC"). Default: IPTypes.PUBLIC + IPTypes.PUBLIC ("PUBLIC"), IPTypes.PRIVATE ("PRIVATE"), + IPTypes.PSC ("PSC"), or IPTypes.SQL_DATA ("SQL_DATA"). + Default: IPTypes.PUBLIC enable_iam_auth (bool): Enables automatic IAM database authentication (Postgres and MySQL) as the default authentication method for all @@ -626,6 +812,15 @@ async def create_async_connector( Must be used with `resolver=DnsResolver` to have any effect. Default: 30 + sql_data_endpoint (str): Endpoint host for SQL Data Service calls. + Default: "sqladmin.googleapis.com". + + sql_data_stream_timeout (int): Timeout in seconds for the SQL Data + Service gRPC stream. Default: 7200. + + resource_exhausted_cooldown_period (float): Cooldown period in seconds + after a ResourceExhausted error. Default: 5.0. + Returns: A Connector instance configured with running event loop. """ @@ -645,4 +840,7 @@ async def create_async_connector( refresh_strategy=refresh_strategy, resolver=resolver, failover_period=failover_period, + sql_data_endpoint=sql_data_endpoint, + sql_data_stream_timeout=sql_data_stream_timeout, + resource_exhausted_cooldown_period=resource_exhausted_cooldown_period, ) diff --git a/google/cloud/sql/connector/enums.py b/google/cloud/sql/connector/enums.py index 4bfb0a44..45134134 100644 --- a/google/cloud/sql/connector/enums.py +++ b/google/cloud/sql/connector/enums.py @@ -41,6 +41,7 @@ class IPTypes(Enum): PUBLIC = "PRIMARY" PRIVATE = "PRIVATE" PSC = "PSC" + SQL_DATA = "SQL_DATA" @classmethod def _missing_(cls, value: object) -> None: @@ -54,6 +55,8 @@ def _from_str(cls, ip_type_str: str) -> IPTypes: """Convert IP type from a str into IPTypes.""" if ip_type_str.upper() == "PUBLIC": ip_type_str = "PRIMARY" + elif ip_type_str.upper() in ("SQLDATA", "SQL_DATA"): + ip_type_str = "SQL_DATA" return cls(ip_type_str.upper()) diff --git a/google/cloud/sql/connector/exceptions.py b/google/cloud/sql/connector/exceptions.py index 4c3d1acb..49d710d7 100644 --- a/google/cloud/sql/connector/exceptions.py +++ b/google/cloud/sql/connector/exceptions.py @@ -14,6 +14,8 @@ limitations under the License. """ +from __future__ import annotations + class ConnectorLoopError(Exception): """ @@ -91,3 +93,34 @@ class ClosedConnectorError(Exception): Exception to be raised when a Connector is closed and connect method is called on it. """ + + +class CloudSQLConnectionError(Exception): + """ + Exception to be raised when a connection cannot be established to a Cloud SQL instance. + """ + + +class ResourceExhaustedError(CloudSQLConnectionError): + """ + Exception to be raised when a connection cannot be established because + the SQL Data Service is busy / in cooldown due to resource exhaustion. + """ + + def __init__( + self, + message: str, + connection_name: str | None = None, + raw_error: Exception | None = None, + ) -> None: + self.message = message + self.connection_name = connection_name + self.raw_error = raw_error + super().__init__( + f"[{connection_name}] {message}: {raw_error}" + if connection_name and raw_error + else f"[{connection_name}] {message}" + if connection_name + else message + ) + diff --git a/google/cloud/sql/connector/monitored_cache.py b/google/cloud/sql/connector/monitored_cache.py index d2ba2290..05d04b0f 100644 --- a/google/cloud/sql/connector/monitored_cache.py +++ b/google/cloud/sql/connector/monitored_cache.py @@ -16,7 +16,7 @@ import asyncio import logging -import ssl +import socket from typing import Any, Callable import aiohttp @@ -44,7 +44,7 @@ def __init__( self.resolver = resolver self.cache = cache self.domain_name_ticker: asyncio.Task | None = None - self.sockets: list[ssl.SSLSocket] = [] + self.sockets: list[socket.socket] = [] # If domain name is configured for instance and failover period is set, # poll for DNS record changes. @@ -77,11 +77,11 @@ def _purge_closed_sockets(self) -> None: list of sockets. """ open_sockets = [] - for socket in self.sockets: + for sock in self.sockets: # Check fileno for if socket is closed. Will return # -1 on failure, which will be used to signal socket closed. - if socket.fileno() != -1: - open_sockets.append(socket) + if sock.fileno() != -1: + open_sockets.append(sock) self.sockets = open_sockets async def _check_domain_name(self) -> None: @@ -159,11 +159,11 @@ async def close(self) -> None: await self.cache.close() # Close any still open sockets - for socket in self.sockets: + for sock in self.sockets: # Check fileno for if socket is closed. Will return # -1 on failure, which will be used to signal socket closed. - if socket.fileno() != -1: - socket.close() + if sock.fileno() != -1: + sock.close() async def ticker(interval: int, function: Callable, *args: Any, **kwargs: Any) -> None: diff --git a/google/cloud/sql/connector/sqldata_client.py b/google/cloud/sql/connector/sqldata_client.py new file mode 100644 index 00000000..acf6287f --- /dev/null +++ b/google/cloud/sql/connector/sqldata_client.py @@ -0,0 +1,620 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import errno +import io +import logging +import queue +import socket +import threading +from typing import Any, Callable + +from google.api_core.client_options import ClientOptions +from google.api_core.exceptions import ResourceExhausted +from google.auth.credentials import Credentials +import grpc + +from google.cloud import sqladmin_v1beta4 +from google.cloud.sql.connector.enums import IPTypes +from google.cloud.sql.connector.exceptions import CloudSQLIPTypeError + +SERVER_PROXY_PORT = 3307 +_EOF_SENTINEL = object() +_STREAM_EOF = object() + +logger = logging.getLogger(__name__) + + +def is_resource_exhausted_error(err: Exception) -> bool: + """Checks whether an exception represents a RESOURCE_EXHAUSTED error.""" + if isinstance(err, ResourceExhausted): + return True + if isinstance(err, grpc.RpcError): + try: + return err.code() == grpc.StatusCode.RESOURCE_EXHAUSTED + except Exception: # noqa: BLE001, S110 + pass + if hasattr(err, "code") and callable(err.code): + try: + return err.code() == grpc.StatusCode.RESOURCE_EXHAUSTED + except Exception: # noqa: BLE001, S110 + pass + cause = getattr(err, "__cause__", None) or getattr(err, "__context__", None) + if isinstance(cause, Exception) and cause is not err: + return is_resource_exhausted_error(cause) + return False + + + +class _RequestQueue: + """Thread-safe queue iterator feeding requests to synchronous gRPC stream.""" + + def __init__(self) -> None: + self._queue: queue.Queue = queue.Queue(maxsize=1024) + self._closed = False + self._lock = threading.Lock() + + def put(self, item: Any) -> None: + with self._lock: + if self._closed: + raise BrokenPipeError(errno.EPIPE, "Stream request queue is closed") + self._queue.put(item) + + def close(self) -> None: + with self._lock: + if self._closed: + return + self._closed = True + self._queue.put(_STREAM_EOF) + + def __iter__(self) -> _RequestQueue: + return self + + def __next__(self) -> Any: + item = self._queue.get() + if item is _STREAM_EOF: + raise StopIteration + return item + + +class SqlDataRawIO(io.RawIOBase): + """RawIO wrapper around SqlDataSocket to support makefile().""" + + def __init__(self, sock: SqlDataSocket) -> None: + self._sock = sock + + def readable(self) -> bool: + return True + + def writable(self) -> bool: + return True + + def seekable(self) -> bool: + return False + + def readinto(self, b: Any) -> int: + return self._sock.recv_into(b) + + def write(self, b: Any) -> int: + self._sock.sendall(b) + return len(b) + + def close(self) -> None: + if not self.closed: + super().close() + self._sock.close() + + +class SqlDataSocket(socket.socket): + """Direct in-process socket adapter connected to a synchronous gRPC SqlData stream. + + Provides full socket interface compatibility for synchronous database drivers + (pg8000, pymysql, pytds) while avoiding local TCP loopback and asyncio scheduling. + """ + + def __init__( + self, + request_queue: _RequestQueue, + response_stream: Any, + grpc_client: sqladmin_v1beta4.SqlDataServiceClient | None = None, + timeout: float | None = None, + on_close: Callable[[], None] | None = None, + on_success: Callable[[], None] | None = None, + on_resource_exhausted: Callable[[Exception], None] | None = None, + fallback_fn: Callable[[], socket.socket] | None = None, + ) -> None: + super().__init__(socket.AF_INET, socket.SOCK_STREAM) + self._request_queue = request_queue + self._response_stream = response_stream + self._grpc_client = grpc_client + self._timeout = timeout + self._on_close = on_close + self._on_success = on_success + self._on_resource_exhausted = on_resource_exhausted + self._fallback_fn = fallback_fn + + self._read_queue: queue.Queue = queue.Queue(maxsize=1024) + self._read_buf = b"" + self._read_offset = 0 + self._closed = False + self._first_read_done = False + self._write_buffer = bytearray() + self._direct_sock: socket.socket | None = None + self._error: Exception | None = None + self._close_lock = threading.Lock() + + # Start background stream consumer thread + self._reader_thread = threading.Thread( + target=self._reader_loop, daemon=True, name="sqldata-reader" + ) + self._reader_thread.start() + + def _reader_loop(self) -> None: + try: + for resp in self._response_stream: + if self._closed: + break + if not self._first_read_done: + self._first_read_done = True + if self._on_success: + self._on_success() + raw_pb = ( + sqladmin_v1beta4.StreamSqlDataResponse.pb(resp) + if hasattr(sqladmin_v1beta4.StreamSqlDataResponse, "pb") + and isinstance(resp, sqladmin_v1beta4.StreamSqlDataResponse) + else resp + ) + which = ( + raw_pb.WhichOneof("message") + if hasattr(raw_pb, "WhichOneof") + else None + ) + if which == "data" or (which is None and hasattr(resp, "data") and resp.data): + data = resp.data.data + if data: + self._read_queue.put(data) + elif which == "session_metadata" or ( + which is None and hasattr(resp, "session_metadata") and resp.session_metadata + ): + logger.debug("Received SessionMetadata") + elif which == "terminate_session" or ( + which is None and hasattr(resp, "terminate_session") and resp.terminate_session + ): + logger.debug("Received TerminateSession from server") + self._closed = True + break + except Exception as e: # noqa: BLE001 + if not self._closed: + logger.debug(f"gRPC sync stream reader encountered: {e}") + self._error = e + if is_resource_exhausted_error(e) and self._on_resource_exhausted: + self._on_resource_exhausted(e) + finally: + self._read_queue.put(_EOF_SENTINEL) + + def sendall( # type: ignore[override] + self, data: Any, flags: int = 0 + ) -> None: + if self._direct_sock is not None: + self._direct_sock.sendall(data, flags) + return + if self._closed: + raise BrokenPipeError(errno.EPIPE, "Socket is closed") + if not data: + return + data_bytes = bytes(data) if not isinstance(data, bytes) else data + if not self._first_read_done: + self._write_buffer.extend(data_bytes) + packet = sqladmin_v1beta4.DataPacket(data=data_bytes) + req = sqladmin_v1beta4.StreamSqlDataRequest(data=packet) + self._request_queue.put(req) + + + def send( # type: ignore[override] + self, data: Any, flags: int = 0 + ) -> int: + if self._direct_sock is not None: + return self._direct_sock.send(data, flags) + self.sendall(data, flags) + return len(data) + + def recv(self, bufsize: int, flags: int = 0) -> bytes: + if self._direct_sock is not None: + return self._direct_sock.recv(bufsize, flags) + + if bufsize <= 0: + return b"" + + # Return from buffered chunk if available + if self._read_offset < len(self._read_buf): + remaining = len(self._read_buf) - self._read_offset + to_copy = min(bufsize, remaining) + chunk = self._read_buf[self._read_offset : self._read_offset + to_copy] + self._read_offset += to_copy + if self._read_offset >= len(self._read_buf): + self._read_buf = b"" + self._read_offset = 0 + return chunk + + # Pull from incoming queue + try: + item = self._read_queue.get(block=True, timeout=self._timeout) + except queue.Empty: + if self._closed: + return b"" + raise socket.timeout("timed out") + + if item is _EOF_SENTINEL: + if self._error is not None: + if ( + not self._first_read_done + and self._fallback_fn is not None + and not is_resource_exhausted_error(self._error) + ): + logger.info( + f"SQL Data Service returned error before first read: {self._error}. " + "Triggering transparent fallback to direct TLS." + ) + try: + self._direct_sock = self._fallback_fn() + if self._write_buffer: + self._direct_sock.sendall(bytes(self._write_buffer)) + self._first_read_done = True + self._write_buffer.clear() + return self._direct_sock.recv(bufsize, flags) + except Exception as fallback_err: + raise fallback_err from self._error + + raise OSError( + errno.ECONNRESET, f"Connection error: {self._error}" + ) from self._error + return b"" + + self._first_read_done = True + self._write_buffer.clear() + if len(item) <= bufsize: + return item + + self._read_buf = item + self._read_offset = bufsize + return item[:bufsize] + + def recv_into(self, buffer: Any, nbytes: int = 0, flags: int = 0) -> int: + if self._direct_sock is not None: + return self._direct_sock.recv_into(buffer, nbytes, flags) + target_len = len(buffer) if nbytes == 0 else min(nbytes, len(buffer)) + if target_len <= 0: + return 0 + data = self.recv(target_len, flags) + n = len(data) + buffer[:n] = data + return n + + def makefile( # type: ignore[override] + self, + mode: str = "r", + buffering: int = -1, + encoding: str | None = None, + errors: str | None = None, + newline: str | None = None, + ) -> Any: + if self._direct_sock is not None: + return self._direct_sock.makefile( # type: ignore[call-overload] + mode=mode, + buffering=buffering, + encoding=encoding, + errors=errors, + newline=newline, + ) + raw = SqlDataRawIO(self) + if buffering == 0: + return raw + + reading = "r" in mode or "+" in mode + writing = "w" in mode or "a" in mode or "+" in mode + binary = "b" in mode + + buf: Any + if reading and writing: + buf = io.BufferedRWPair(raw, raw) + elif reading: + buffer_size = io.DEFAULT_BUFFER_SIZE if buffering <= 0 else buffering + buf = io.BufferedReader(raw, buffer_size=buffer_size) + elif writing: + buffer_size = io.DEFAULT_BUFFER_SIZE if buffering <= 0 else buffering + buf = io.BufferedWriter(raw, buffer_size=buffer_size) + else: + buffer_size = io.DEFAULT_BUFFER_SIZE if buffering <= 0 else buffering + buf = io.BufferedReader(raw, buffer_size=buffer_size) + + if binary: + return buf + return io.TextIOWrapper( + buf, encoding=encoding, errors=errors, newline=newline + ) + + def settimeout(self, value: float | None) -> None: + if value is not None and value < 0: + raise ValueError("Timeout value must be non-negative") + self._timeout = value + if self._direct_sock is not None: + self._direct_sock.settimeout(value) + + def gettimeout(self) -> float | None: + if self._direct_sock is not None: + return self._direct_sock.gettimeout() + return self._timeout + + def setblocking(self, flag: bool) -> None: + self._timeout = None if flag else 0.0 + if self._direct_sock is not None: + self._direct_sock.setblocking(flag) + + def connect(self, *args: Any, **kwargs: Any) -> None: + # Already connected, no-op for driver compatibility (e.g. pymysql) + pass + + def connect_ex(self, *args: Any, **kwargs: Any) -> int: + return 0 + + def setsockopt(self, *args: Any, **kwargs: Any) -> None: + if self._direct_sock is not None: + self._direct_sock.setsockopt(*args, **kwargs) + + def getsockopt(self, *args: Any, **kwargs: Any) -> Any: # type: ignore[override] + if self._direct_sock is not None: + return self._direct_sock.getsockopt(*args, **kwargs) + return 0 + + def getsockname(self) -> tuple[str, int]: + if self._direct_sock is not None: + return self._direct_sock.getsockname() + return ("127.0.0.1", SERVER_PROXY_PORT) + + def getpeername(self) -> tuple[str, int]: + if self._direct_sock is not None: + return self._direct_sock.getpeername() + return ("127.0.0.1", SERVER_PROXY_PORT) + + def shutdown(self, how: int = socket.SHUT_RDWR) -> None: + self.close() + + def close(self) -> None: + with self._close_lock: + if self._closed: + return + self._closed = True + + if self._direct_sock is not None: + try: + self._direct_sock.close() + except Exception: # noqa: BLE001, S110 + pass + + self._request_queue.close() + self._read_queue.put(_EOF_SENTINEL) + try: + if hasattr(self._response_stream, "cancel"): + self._response_stream.cancel() + except Exception: # noqa: BLE001, S110 + pass + try: + if ( + self._grpc_client is not None + and hasattr(self._grpc_client, "transport") + and hasattr(self._grpc_client.transport, "close") + ): + self._grpc_client.transport.close() + except Exception: # noqa: BLE001, S110 + pass + try: + super().close() + except Exception: # noqa: BLE001, S110 + pass + + if self._on_close: + try: + self._on_close() + except Exception: # noqa: BLE001, S110 + pass + + +class SqlDataClient: + """Client that establishes direct synchronous gRPC SqlDataService connections.""" + + def __init__( + self, + endpoint: str, + credentials: Credentials, + quota_project: str | None = None, + timeout: float | None = None, + ) -> None: + self._endpoint = endpoint + self._credentials = credentials + self._quota_project = quota_project + self._timeout = timeout + self._active_sockets: set[SqlDataSocket] = set() + self._on_close_callbacks: list[Callable[[], None]] = [] + + async def connect( + self, + instance_connection_name: str, + region: str, + project: str, + get_conn_info: Callable[[], Any], + enable_iam_auth: bool, + on_fallback: Callable[[str], None], + is_fallback_cached: Callable[[str], bool], + on_resource_exhausted: Callable[[Exception], None] | None = None, + on_success: Callable[[], None] | None = None, + connect_timeout: float = 30.0, + ) -> socket.socket: + """Connects via synchronous gRPC and returns a SqlDataSocket or direct TLS fallback socket.""" + use_fallback = is_fallback_cached(instance_connection_name) + + async def get_direct_socket_sync_factory() -> Callable[[], socket.socket]: + conn_info = await get_conn_info() + targets: list[str] = [] + for t in [IPTypes.PRIVATE, IPTypes.PSC, IPTypes.PUBLIC]: + try: + targets.extend(conn_info.get_preferred_ips(t)) + except CloudSQLIPTypeError as e: + logger.debug(f"IP type {t} not available: {e}") + continue + if not targets: + raise ValueError( + "Cannot fallback to direct connection: no IP address available." + ) + ssl_context = await conn_info.create_ssl_context(enable_iam_auth) + + def create_direct_sock() -> socket.socket: + last_ex: Exception | None = None + for target_ip in targets: + logger.debug(f"Direct TLS connecting to {target_ip}:{SERVER_PROXY_PORT}") + try: + raw_sock = socket.create_connection( + (target_ip, SERVER_PROXY_PORT), timeout=connect_timeout + ) + ssl_sock = ssl_context.wrap_socket( + raw_sock, server_hostname=target_ip + ) + return ssl_sock + except Exception as e: # noqa: BLE001 + logger.debug(f"Direct TLS connection to {target_ip} failed: {e}") + last_ex = e + if last_ex: + raise last_ex + raise ValueError( + "Cannot fallback to direct connection: no IP address available." + ) + + return create_direct_sock + + if use_fallback: + logger.debug("Using cached fallback direct TLS connection") + create_direct = await get_direct_socket_sync_factory() + return create_direct() + + # Connect synchronous gRPC stream via GAPIC client + endpoint = self._endpoint.removeprefix("https://").removeprefix("http://") + client_options = ClientOptions( + api_endpoint=endpoint, + quota_project_id=self._quota_project, + ) + grpc_client = sqladmin_v1beta4.SqlDataServiceClient( + credentials=self._credentials, + client_options=client_options, + ) + + instance_id = ( + f"projects/{project}/instances/{instance_connection_name.split(':')[-1]}" + ) + location_id = f"locations/{region}" + + metadata = [ + ( + "x-goog-request-params", + f"instance_id={instance_id}&location_id={location_id}", + ) + ] + + create_direct_sock_fn: Callable[[], socket.socket] | None = None + try: + create_direct_sock_fn = await get_direct_socket_sync_factory() + except Exception as e: # noqa: BLE001 + logger.debug(f"Could not prepare direct fallback factory: {e}") + + def fallback_fn() -> socket.socket: + on_fallback(instance_connection_name) + if create_direct_sock_fn is None: + raise ValueError("No direct fallback connection factory available.") + return create_direct_sock_fn() + + try: + request_queue = _RequestQueue() + start_session = sqladmin_v1beta4.StartSession( + instance_id=instance_id, location_id=location_id + ) + req = sqladmin_v1beta4.StreamSqlDataRequest( + start_session=start_session + ) + request_queue.put(req) + + response_stream: Any + if hasattr(grpc_client, "transport") and hasattr( + grpc_client.transport, "stream_sql_data" + ): + response_stream = grpc_client.transport.stream_sql_data( # type: ignore[call-arg] + request_queue, # type: ignore[arg-type] + metadata=metadata, + timeout=self._timeout, + ) + else: + response_stream = grpc_client.stream_sql_data( # type: ignore[arg-type] + requests=request_queue, metadata=metadata, timeout=self._timeout + ) + + sock = SqlDataSocket( + request_queue=request_queue, + response_stream=response_stream, + grpc_client=grpc_client, + timeout=connect_timeout, + on_close=lambda: self._active_sockets.discard(sock), + on_success=on_success, + on_resource_exhausted=on_resource_exhausted, + fallback_fn=fallback_fn, + ) + self._active_sockets.add(sock) + return sock + + except Exception as e: + logger.debug(f"Sync gRPC connection attempt failed: {e}") + try: + if hasattr(grpc_client, "transport") and hasattr(grpc_client.transport, "close"): + grpc_client.transport.close() + except Exception: # noqa: BLE001, S110 + pass + + if is_resource_exhausted_error(e): + if on_resource_exhausted: + on_resource_exhausted(e) + raise + + # Fallback to direct TLS on connection failure + logger.info( + f"SQL Data Service connection failed for {instance_connection_name}. " + "Falling back to direct TLS connection." + ) + on_fallback(instance_connection_name) + if create_direct_sock_fn: + return create_direct_sock_fn() + create_direct = await get_direct_socket_sync_factory() + return create_direct() + + + async def close(self) -> None: + """Closes all active sockets created by this client.""" + for sock in list(self._active_sockets): + try: + sock.close() + except Exception: # noqa: BLE001, S110 + pass + self._active_sockets.clear() + for cb in self._on_close_callbacks: + try: + cb() + except Exception: # noqa: BLE001, S110 + pass diff --git a/pyproject.toml b/pyproject.toml index 93ed5eaf..13f83b85 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -46,10 +46,9 @@ dependencies = [ "dnspython>=2.0.0", "Requests", "google-auth>=2.28.0", - "grpcio", - "protobuf", - "googleapis-common-protos", + "google-cloud-sql>=0.1.1", ] + dynamic = ["version"] [project.urls] @@ -97,3 +96,17 @@ single-line-exclusions = ["typing"] [tool.ruff.format] quote-style = "double" + +[tool.coverage.report] +exclude_lines = [ + "pragma: no cover", + "def __repr__", + "raise AssertionError", + "raise NotImplementedError", + "if __name__ == .__main__.:", + "@abc.abstractmethod", + "if TYPE_CHECKING:", +] + + + diff --git a/requirements.txt b/requirements.txt index 27889567..238a1625 100644 --- a/requirements.txt +++ b/requirements.txt @@ -2,7 +2,9 @@ certifi==2026.7.22 cryptography==50.0.0 filelock==3.32.3 google-auth==2.56.3 +google-cloud-sql==0.1.1 idna==3.19 + packaging==26.3 pip==26.2.1 PyMySQL==1.2.0 diff --git a/tests/system/test_asyncpg_connection.py b/tests/system/test_asyncpg_connection.py index 72f4dabb..269a6cc2 100644 --- a/tests/system/test_asyncpg_connection.py +++ b/tests/system/test_asyncpg_connection.py @@ -283,3 +283,5 @@ async def test_lazy_connection_with_asyncpg() -> None: assert res[0][0] == 1 await connector.close_async() + + diff --git a/tests/system/test_pg8000_connection.py b/tests/system/test_pg8000_connection.py index 4a2464cb..1458989b 100644 --- a/tests/system/test_pg8000_connection.py +++ b/tests/system/test_pg8000_connection.py @@ -19,6 +19,7 @@ import os # [START cloud_sql_connector_postgres_pg8000] +import pytest import sqlalchemy from google.cloud.sql.connector import Connector @@ -209,3 +210,58 @@ def test_MCP_pg8000_connection() -> None: curr_time = time[0] assert type(curr_time) is datetime connector.close() + + +def test_AIDE_pg8000_connection() -> None: + """Basic test to get time from database using AIDE instance.""" + if "POSTGRES_AIDE_CONNECTION_NAME" not in os.environ: + pytest.skip("POSTGRES_AIDE_CONNECTION_NAME not set") + inst_conn_name = os.environ["POSTGRES_AIDE_CONNECTION_NAME"] + user = os.environ.get("POSTGRES_AIDE_USER", os.environ.get("POSTGRES_USER", "postgres")) + password = os.environ.get("POSTGRES_AIDE_PASS", os.environ.get("POSTGRES_PASS", "")) + db = os.environ.get("POSTGRES_AIDE_DB", os.environ.get("POSTGRES_DB", "postgres")) + + engine, connector = create_sqlalchemy_engine( + inst_conn_name, + user, + password, + db, + ip_type="sqldata", + ) + with engine.connect() as conn: + time = conn.execute(sqlalchemy.text("SELECT NOW()")).fetchone() + conn.commit() + curr_time = time[0] + assert type(curr_time) is datetime + connector.close() + + +def test_sqldata_fallback_pg8000_connection() -> None: + """Test connecting to a non-AIDE instance with ip_type='sqldata'. + + The server returns FAILED_PRECONDITION for standard (non-Developer Edition) instances, + and the connector falls back to connecting via public IP. + """ + if "POSTGRES_FALLBACK_CONNECTION_NAME" not in os.environ: + pytest.skip("POSTGRES_FALLBACK_CONNECTION_NAME not set") + inst_conn_name = os.environ["POSTGRES_FALLBACK_CONNECTION_NAME"] + user = os.environ.get("POSTGRES_FALLBACK_USER", os.environ.get("POSTGRES_USER", "postgres")) + password = os.environ.get("POSTGRES_FALLBACK_PASS", os.environ.get("POSTGRES_PASS", "")) + db = os.environ.get("POSTGRES_FALLBACK_DB", os.environ.get("POSTGRES_DB", "postgres")) + + engine, connector = create_sqlalchemy_engine( + inst_conn_name, + user, + password, + db, + ip_type="sqldata", + ) + with engine.connect() as conn: + time = conn.execute(sqlalchemy.text("SELECT NOW()")).fetchone() + conn.commit() + curr_time = time[0] + assert type(curr_time) is datetime + connector.close() + + + diff --git a/tests/unit/test_client.py b/tests/unit/test_client.py index bdb3e568..76c5e19e 100644 --- a/tests/unit/test_client.py +++ b/tests/unit/test_client.py @@ -298,7 +298,33 @@ async def test_get_ephemeral_error_parsing_json( await client.close() +async def test_get_ephemeral_missing_cert_key( + fake_credentials: Credentials, +) -> None: + """ + Test that KeyError is raised and logged when ephemeralCert is missing. + """ + client = CloudSQLClient( + sqladmin_api_endpoint="https://sqladmin.googleapis.com", + quota_project=None, + credentials=fake_credentials, + ) + post_url = "https://sqladmin.googleapis.com/sql/v1beta4/projects/my-project/instances/my-instance:generateEphemeralCert" + resp_body = {} # missing ephemeralCert + with aioresponses() as mocked: + mocked.post( + post_url, + status=200, + payload=resp_body, + repeat=True, + ) + with pytest.raises(KeyError): + await client._get_ephemeral("my-project", "my-instance", "my-key") + await client.close() + + @pytest.mark.asyncio + async def test_get_metadata_multiple_psc_dns_sorted(fake_client: CloudSQLClient) -> None: """ Test _get_metadata returns successfully with multiple PSC IP types sorted. diff --git a/tests/unit/test_connector.py b/tests/unit/test_connector.py index 9b621a69..4e49643d 100644 --- a/tests/unit/test_connector.py +++ b/tests/unit/test_connector.py @@ -17,6 +17,7 @@ import asyncio import os +import socket from threading import Thread from unittest.mock import AsyncMock from unittest.mock import MagicMock @@ -39,6 +40,7 @@ from google.cloud.sql.connector.instance import RefreshAheadCache from google.cloud.sql.connector.monitored_cache import MonitoredCache from google.cloud.sql.connector.resolver import DnsResolver +from google.cloud.sql.connector.sqldata_client import SqlDataClient @pytest.mark.asyncio @@ -245,7 +247,7 @@ def test_Connector_Init_bad_ip_type(fake_credentials: Credentials) -> None: assert ( exc_info.value.args[0] == f"Incorrect value for ip_type, got '{bad_ip_type.upper()}'. " - "Want one of: 'PRIMARY', 'PRIVATE', 'PSC', 'PUBLIC'." + "Want one of: 'PRIMARY', 'PRIVATE', 'PSC', 'SQL_DATA', 'PUBLIC'." ) @@ -268,7 +270,7 @@ def test_Connector_connect_bad_ip_type( assert ( exc_info.value.args[0] == f"Incorrect value for ip_type, got '{bad_ip_type.upper()}'. " - "Want one of: 'PRIMARY', 'PRIVATE', 'PSC', 'PUBLIC'." + "Want one of: 'PRIMARY', 'PRIVATE', 'PSC', 'SQL_DATA', 'PUBLIC'." ) @@ -1064,5 +1066,506 @@ async def test_Connector_connect_async_psycopg( mock_connect.assert_called_once() +def test_Connector_Init_sqldata_options(fake_credentials: Credentials) -> None: + """Test that Connector initializes with custom SQL data endpoint and timeout.""" + with Connector( + credentials=fake_credentials, + sql_data_endpoint="custom.sqladmin.googleapis.com", + sql_data_stream_timeout=3600, + ) as connector: + assert connector._sql_data_endpoint == "custom.sqladmin.googleapis.com" + assert connector._sql_data_stream_timeout == 3600 + + +@pytest.mark.asyncio +async def test_create_async_connector_sqldata_options( + fake_credentials: Credentials, +) -> None: + """Test that create_async_connector properly forwards SQL data options.""" + connector = await create_async_connector( + credentials=fake_credentials, + sql_data_endpoint="custom.sqladmin.googleapis.com", + sql_data_stream_timeout=1800, + ) + assert connector._sql_data_endpoint == "custom.sqladmin.googleapis.com" + assert connector._sql_data_stream_timeout == 1800 + await connector.close_async() + + +@pytest.mark.asyncio +async def test_Connector_connect_async_sqldata_iam_auth( + fake_credentials: Credentials, + fake_client: CloudSQLClient, +) -> None: + """Test that connect_async with SQL_DATA and IAM auth properly maps driver engine without KeyError.""" + connect_string = "test-project:test-region:test-instance" + async with Connector( + credentials=fake_credentials, + loop=asyncio.get_running_loop(), + ip_type=IPTypes.SQL_DATA, + ) as connector: + connector._client = fake_client + + with patch("google.cloud.sql.connector.connector.SqlDataClient") as mock_sqldata_cls: + mock_client_instance = MagicMock() + mock_sock = MagicMock(spec=socket.socket) + mock_client_instance.connect = AsyncMock(return_value=mock_sock) + mock_client_instance.close = AsyncMock() + mock_sqldata_cls.return_value = mock_client_instance + + with patch("google.cloud.sql.connector.pg8000.connect") as mock_connect: + mock_connect.return_value = True + + connection = await connector.connect_async( + connect_string, + "pg8000", + user="test-sa@test-project.iam.gserviceaccount.com", + db="my-db", + enable_iam_auth=True, + ) + assert connection is True + # Verify IAM user was formatted and passed without error + assert mock_connect.called + _, kwargs = mock_connect.call_args + assert kwargs["user"] == "test-sa@test-project.iam" + + +def test_sqldata_client_init(fake_credentials: Credentials) -> None: + """Test that SqlDataClient initializes with expected properties.""" + client = SqlDataClient( + endpoint="custom.sqladmin.googleapis.com", + credentials=fake_credentials, + quota_project="test-quota-project", + timeout=3600, + ) + assert client._endpoint == "custom.sqladmin.googleapis.com" + assert client._credentials == fake_credentials + assert client._quota_project == "test-quota-project" + assert client._timeout == 3600 + assert len(client._active_sockets) == 0 + + +@pytest.mark.asyncio +async def test_sqldata_client_close(fake_credentials: Credentials) -> None: + """Test that SqlDataClient.close cleanly closes active sockets.""" + client = SqlDataClient( + endpoint="sqladmin.googleapis.com", + credentials=fake_credentials, + ) + mock_sock = MagicMock() + client._active_sockets.add(mock_sock) + + callback_called = False + + def on_close() -> None: + nonlocal callback_called + callback_called = True + + client._on_close_callbacks.append(on_close) + + await client.close() + + assert mock_sock.close.called + assert len(client._active_sockets) == 0 + assert callback_called + + +def test_Connector_Init_resource_exhausted_options( + fake_credentials: Credentials, +) -> None: + """Test that Connector initializes with resource_exhausted_cooldown_period.""" + with Connector( + credentials=fake_credentials, + resource_exhausted_cooldown_period=10.0, + ) as connector: + assert connector._resource_exhausted_cooldown_period == 10.0 + + +@pytest.mark.asyncio +async def test_create_async_connector_resource_exhausted_options( + fake_credentials: Credentials, +) -> None: + """Test that create_async_connector forwards resource_exhausted_cooldown_period.""" + connector = await create_async_connector( + credentials=fake_credentials, + resource_exhausted_cooldown_period=8.5, + ) + assert connector._resource_exhausted_cooldown_period == 8.5 + await connector.close_async() + + +def test_cooldown_backoff_calculation() -> None: + """Test exponential backoff with jitter calculation.""" + from google.cloud.sql.connector.connector import _cooldown_backoff + + base = 5.0 + for attempt in range(1, 6): + backoff = _cooldown_backoff(base, attempt) + # 1.618^(attempt-1) <= multiplier <= 1.618^attempt + min_expected = base * (1.618 ** (attempt - 1)) + max_expected = base * (1.618**attempt) + assert min_expected <= backoff <= max_expected + + +def test_is_resource_exhausted_error_helper() -> None: + """Test is_resource_exhausted_error helper with various exception types.""" + import grpc + + from google.cloud.sql.connector.sqldata_client import is_resource_exhausted_error + + class MockRpcError(Exception): + def __init__(self, code): + self._code = code + + def code(self): + return self._code + + assert is_resource_exhausted_error( + MockRpcError(grpc.StatusCode.RESOURCE_EXHAUSTED) + ) + assert not is_resource_exhausted_error( + MockRpcError(grpc.StatusCode.FAILED_PRECONDITION) + ) + assert not is_resource_exhausted_error(Exception("other error")) + + # Test wrapped cause + wrapped = Exception("wrapper error") + wrapped.__cause__ = MockRpcError(grpc.StatusCode.RESOURCE_EXHAUSTED) + assert is_resource_exhausted_error(wrapped) + + +def test_SqlDataConnState_methods() -> None: + """Test SqlDataConnState state transitions and helper methods.""" + import time + + from google.cloud.sql.connector.connector import SqlDataConnState + + state = SqlDataConnState() + assert state.allowed is True + assert state.is_cooldown_active() is False + + err = Exception("resource busy") + backoff = state.record_exhausted(err, base_cooldown=2.0) + assert state.backoff_counter == 1 + assert state.last_err is err + assert state.cooldown_until is not None + assert state.cooldown_until > time.time() + assert state.is_cooldown_active() is True + assert backoff > 0 + + state.record_success() + assert state.backoff_counter == 0 + assert state.cooldown_until is None + assert state.last_err is None + assert state.is_cooldown_active() is False + + state.record_fallback() + assert state.allowed is False + assert state.is_cooldown_active() is False + + +@pytest.mark.asyncio +async def test_ResourceExhausted_cooldown_blocks_connection( + fake_credentials: Credentials, + fake_client: CloudSQLClient, +) -> None: + """Test that active cooldown raises ResourceExhaustedError without connecting.""" + import time + + from google.cloud.sql.connector.connector import SqlDataConnState + from google.cloud.sql.connector.exceptions import ResourceExhaustedError + + connect_string = "proj:reg:inst" + async with Connector( + credentials=fake_credentials, + loop=asyncio.get_running_loop(), + ip_type=IPTypes.SQL_DATA, + resource_exhausted_cooldown_period=2.0, + ) as connector: + connector._client = fake_client + + # Manually set state to cooldown active + state = SqlDataConnState() + state.cooldown_until = time.time() + 10.0 + state.backoff_counter = 1 + state.last_err = Exception("resource busy") + connector._sql_data_conn_state[connect_string] = state + + with ( + patch( + "google.cloud.sql.connector.connector.SqlDataClient" + ) as mock_sqldata_cls, + pytest.raises(ResourceExhaustedError) as exc_info, + ): + await connector.connect_async( + connect_string, + "pg8000", + user="test-user", + db="test-db", + ) + assert "cooldown active" in str(exc_info.value) + assert not mock_sqldata_cls.called + + +@pytest.mark.asyncio +async def test_ResourceExhausted_callbacks_lifecycle( + fake_credentials: Credentials, + fake_client: CloudSQLClient, +) -> None: + """Test that on_resource_exhausted and on_success callbacks properly update state.""" + import time + + from google.cloud.sql.connector.exceptions import ResourceExhaustedError + + connect_string = "proj:reg:inst" + async with Connector( + credentials=fake_credentials, + loop=asyncio.get_running_loop(), + ip_type=IPTypes.SQL_DATA, + resource_exhausted_cooldown_period=0.5, + ) as connector: + connector._client = fake_client + + captured_on_resource_exhausted = None + captured_on_success = None + + mock_sqldata_instance = MagicMock() + + async def mock_connect(**kwargs): + nonlocal captured_on_resource_exhausted, captured_on_success + captured_on_resource_exhausted = kwargs.get("on_resource_exhausted") + captured_on_success = kwargs.get("on_success") + return MagicMock(spec=socket.socket) + + mock_sqldata_instance.connect = AsyncMock(side_effect=mock_connect) + mock_sqldata_instance.close = AsyncMock() + + with ( + patch( + "google.cloud.sql.connector.connector.SqlDataClient", + return_value=mock_sqldata_instance, + ), + patch("google.cloud.sql.connector.pg8000.connect", return_value=True), + ): + # 1. Connect and trigger on_resource_exhausted + await connector.connect_async( + connect_string, + "pg8000", + user="test-user", + db="test-db", + ) + assert captured_on_resource_exhausted is not None + assert captured_on_success is not None + + state = connector._sql_data_conn_state[connect_string] + assert state.backoff_counter == 0 + assert state.cooldown_until is None + + # Trigger resource exhausted + captured_on_resource_exhausted(Exception("resource exhausted")) + assert state.backoff_counter == 1 + assert state.cooldown_until is not None + assert state.cooldown_until > time.time() + + # Second connect attempt during cooldown fails with ResourceExhaustedError + with pytest.raises(ResourceExhaustedError): + await connector.connect_async( + connect_string, + "pg8000", + user="test-user", + db="test-db", + ) + + # Reset via on_success + captured_on_success() + assert state.backoff_counter == 0 + assert state.cooldown_until is None + assert state.last_err is None + + +@pytest.mark.asyncio +async def test_sqldata_fallback_ip_order(fake_credentials: Credentials) -> None: + """Test that direct fallback queries IP addresses in PRIVATE, PSC, PUBLIC order.""" + client = SqlDataClient( + endpoint="sqladmin.googleapis.com", + credentials=fake_credentials, + ) + mock_conn_info = MagicMock() + queried_ip_types: list[IPTypes] = [] + + def mock_get_preferred_ips(ip_type: IPTypes): + queried_ip_types.append(ip_type) + if ip_type == IPTypes.PUBLIC: + return ["1.2.3.4"] + from google.cloud.sql.connector.exceptions import CloudSQLIPTypeError + + raise CloudSQLIPTypeError(f"{ip_type} not available") + + mock_conn_info.get_preferred_ips.side_effect = mock_get_preferred_ips + mock_ssl_ctx = MagicMock() + mock_ssl_sock = MagicMock(spec=socket.socket) + mock_ssl_ctx.wrap_socket.return_value = mock_ssl_sock + mock_conn_info.create_ssl_context = AsyncMock(return_value=mock_ssl_ctx) + get_conn_info = AsyncMock(return_value=mock_conn_info) + + mock_raw_sock = MagicMock(spec=socket.socket) + + with patch("socket.create_connection", return_value=mock_raw_sock): + sock = await client.connect( + instance_connection_name="proj:reg:inst", + region="reg", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=MagicMock(), + is_fallback_cached=MagicMock(return_value=True), + ) + + assert sock is mock_ssl_sock + assert queried_ip_types == [IPTypes.PRIVATE, IPTypes.PSC, IPTypes.PUBLIC] + await client.close() + + +@pytest.mark.asyncio +async def test_Connector_connect_async_sqldata_incompatible_driver( + fake_credentials: Credentials, + fake_client: CloudSQLClient, +) -> None: + """Test that connecting with SQL_DATA and an async driver raises IncompatibleDriverError.""" + async with Connector( + credentials=fake_credentials, + loop=asyncio.get_running_loop(), + ip_type=IPTypes.SQL_DATA, + ) as connector: + connector._client = fake_client + + with pytest.raises(IncompatibleDriverError, match="Driver 'asyncpg' is not supported"): + await connector.connect_async( + "test-project:test-region:test-instance", + "asyncpg", + user="my-user", + password="my-pass", + db="my-db", + ) + + +@pytest.mark.asyncio +async def test_Connector_connect_async_sqldata_domain_name_and_error_cleanup( + fake_credentials: Credentials, + fake_client: CloudSQLClient, +) -> None: + """Test that connecting with SQL_DATA and a domain name manages socket cache and cleans up on error.""" + from google.cloud.sql.connector.resolver import DnsResolver + + mock_sqldata_client = MagicMock() + mock_sock = MagicMock() + mock_sqldata_client.connect = AsyncMock(return_value=mock_sock) + mock_sqldata_client.close = AsyncMock() + mock_sqldata_client._on_close_callbacks = [] + + with patch( + "google.cloud.sql.connector.resolver.DnsResolver.resolve", + return_value=ConnectionName( + "test-project", "test-region", "test-instance", "db.example.com" + ), + ), patch( + "google.cloud.sql.connector.connector.SqlDataClient", + return_value=mock_sqldata_client, + ): + async with Connector( + credentials=fake_credentials, + loop=asyncio.get_running_loop(), + ip_type=IPTypes.SQL_DATA, + resolver=DnsResolver, + ) as connector: + connector._client = fake_client + + with patch( + "google.cloud.sql.connector.pg8000.connect", + side_effect=RuntimeError("pg8000 connect failed"), + ), patch.object( + MonitoredCache, "force_refresh", AsyncMock() + ): + with pytest.raises(RuntimeError, match="pg8000 connect failed"): + await connector.connect_async( + "db.example.com", + "pg8000", + user="my-user", + password="my-pass", + db="my-db", + ) + mock_sock.close.assert_called_once() + + + +@pytest.mark.asyncio +async def test_Connector_connect_async_sqldata_fallback_and_callbacks( + fake_credentials: Credentials, + fake_client: CloudSQLClient, +) -> None: + """Test on_fallback, on_success, and is_fallback_cached callbacks for SQL_DATA.""" + mock_sqldata_client = MagicMock() + mock_sock = MagicMock() + mock_sqldata_client.close = AsyncMock() + mock_sqldata_client._on_close_callbacks = [] + + async def fake_connect(**kwargs): + # Trigger on_fallback and verify + on_fallback = kwargs["on_fallback"] + is_fallback_cached = kwargs["is_fallback_cached"] + on_success = kwargs["on_success"] + get_conn_info = kwargs["get_conn_info"] + + # Call get_conn_info + conn_info = await get_conn_info() + assert conn_info is not None + + # Verify initial fallback cache state + assert not is_fallback_cached("test-project:test-region:test-instance") + + # Trigger fallback + on_fallback("test-project:test-region:test-instance") + assert is_fallback_cached("test-project:test-region:test-instance") + + on_success() + return mock_sock + + mock_sqldata_client.connect = AsyncMock(side_effect=fake_connect) + + with patch( + "google.cloud.sql.connector.connector.SqlDataClient", + return_value=mock_sqldata_client, + ): + async with Connector( + credentials=fake_credentials, + loop=asyncio.get_running_loop(), + ip_type=IPTypes.SQL_DATA, + ) as connector: + connector._client = fake_client + + with patch("google.cloud.sql.connector.pg8000.connect", return_value=True): + conn = await connector.connect_async( + "test-project:test-region:test-instance", + "pg8000", + user="my-user", + password="my-pass", + db="my-db", + ) + assert conn is True + + + + +def test_Connector_close_handles_exception(fake_credentials: Credentials) -> None: + """Test Connector.close() safely handles exception if close_future raises.""" + connector = Connector(credentials=fake_credentials) + with patch( + "asyncio.run_coroutine_threadsafe" + ) as mock_run: + mock_future = MagicMock() + mock_future.result.side_effect = TimeoutError("Timed out") + mock_run.return_value = mock_future + # Should log and not raise exception + connector.close() diff --git a/tests/unit/test_instance.py b/tests/unit/test_instance.py index ca87370b..5510c0ce 100644 --- a/tests/unit/test_instance.py +++ b/tests/unit/test_instance.py @@ -28,6 +28,7 @@ from google.cloud.sql.connector.connection_info import ConnectionInfo from google.cloud.sql.connector.connection_name import ConnectionName from google.cloud.sql.connector.exceptions import AutoIAMAuthNotSupported +from google.cloud.sql.connector.exceptions import CloudSQLConnectionError from google.cloud.sql.connector.exceptions import CloudSQLIPTypeError from google.cloud.sql.connector.exceptions import TLSVersionError from google.cloud.sql.connector.instance import RefreshAheadCache @@ -391,3 +392,13 @@ 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_ConnectionInfo_missing_server_ca_cert() -> None: + """Test that create_ssl_context raises CloudSQLConnectionError when server_ca_cert is None.""" + info = ConnectionInfo( + "", "cert", None, b"key", {}, "POSTGRES", datetime.datetime.now(datetime.timezone.utc) + ) + with pytest.raises(CloudSQLConnectionError) as exc_info: + await info.create_ssl_context() + assert "server CA certificate is missing" in str(exc_info.value) diff --git a/tests/unit/test_sqldata_client.py b/tests/unit/test_sqldata_client.py new file mode 100644 index 00000000..4a8ac9c1 --- /dev/null +++ b/tests/unit/test_sqldata_client.py @@ -0,0 +1,1041 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import queue +import socket +import time +from unittest.mock import AsyncMock +from unittest.mock import MagicMock +from unittest.mock import patch + +from google.api_core.exceptions import ResourceExhausted +from google.auth.credentials import Credentials +import grpc +import pytest + +from google.cloud import sqladmin_v1beta4 +from google.cloud.sql.connector.sqldata_client import _RequestQueue +from google.cloud.sql.connector.sqldata_client import is_resource_exhausted_error +from google.cloud.sql.connector.sqldata_client import SqlDataClient +from google.cloud.sql.connector.sqldata_client import SqlDataSocket + + +class MockRpcError(grpc.RpcError): + def __init__(self, code: grpc.StatusCode): + self._code = code + + def code(self) -> grpc.StatusCode: + return self._code + + +def test_is_resource_exhausted_error(): + # Regular exception + assert not is_resource_exhausted_error(ValueError("foo")) + + # ResourceExhausted from google.api_core.exceptions + assert is_resource_exhausted_error(ResourceExhausted("quota exceeded")) + + # RpcError with RESOURCE_EXHAUSTED + mock_err = MockRpcError(grpc.StatusCode.RESOURCE_EXHAUSTED) + assert is_resource_exhausted_error(mock_err) + + # RpcError with other status + mock_err_other = MockRpcError(grpc.StatusCode.UNAVAILABLE) + assert not is_resource_exhausted_error(mock_err_other) + + # Wrapped exception + wrapped = Exception("wrapped") + wrapped.__cause__ = mock_err + assert is_resource_exhausted_error(wrapped) + + +def test_sqldata_socket_send_recv(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + + data_packet = sqladmin_v1beta4.DataPacket(data=b"hello world") + resp1 = sqladmin_v1beta4.StreamSqlDataResponse(data=data_packet) + + def stream_gen(): + yield resp1 + # Block until stream cancelled/closed + time.sleep(1.0) + + mock_response_stream.__iter__.side_effect = stream_gen + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + timeout=2.0, + ) + + # Test sendall + sock.sendall(b"client query") + written_req = next(iter(req_queue)) + assert written_req.data.data == b"client query" + + # Test send + sent_len = sock.send(b"12345") + assert sent_len == 5 + + # Test recv chunked + chunk1 = sock.recv(5) + assert chunk1 == b"hello" + + # Test recv_into + buf = bytearray(5) + n = sock.recv_into(buf) + assert n == 5 + assert bytes(buf) == b" worl" + + # Test makefile + sock._read_queue.put(b"line1\nline2\n") + rfile = sock.makefile("rb") + line = rfile.readline() + assert line == b"dline1\n" or line == b"line1\n" or line.endswith(b"line1\n") + + # Test socket options & no-ops + sock.settimeout(5.0) + assert sock.gettimeout() == 5.0 + sock.setblocking(True) + assert sock.gettimeout() is None + sock.connect(("127.0.0.1", 3307)) + assert sock.connect_ex(("127.0.0.1", 3307)) == 0 + sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) + assert sock.getsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY) == 0 + assert sock.getsockname() == ("127.0.0.1", 3307) + assert sock.getpeername() == ("127.0.0.1", 3307) + + # Test close + sock.close() + assert sock._closed + with pytest.raises(BrokenPipeError): + sock.sendall(b"after close") + + + +def test_is_resource_exhausted_error_code_exception(): + class BrokenError(Exception): + def code(self): + raise RuntimeError("code failed") + + class BrokenRpcError(grpc.RpcError): + def code(self): + raise RuntimeError("code failed") + + err = BrokenError() + assert not is_resource_exhausted_error(err) + assert not is_resource_exhausted_error(BrokenRpcError()) + + # Wrapped with code exception in parent but valid in cause + wrapped = Exception("wrapper") + wrapped.__cause__ = MockRpcError(grpc.StatusCode.RESOURCE_EXHAUSTED) + err.__cause__ = wrapped + assert is_resource_exhausted_error(err) + + + +def test_request_queue_operations(): + q = _RequestQueue() + q.put("item1") + assert next(q) == "item1" + + q.close() + # Second close should be a no-op + q.close() + + # Next after close should raise StopIteration + with pytest.raises(StopIteration): + next(q) + + # Put after close should raise BrokenPipeError + with pytest.raises(BrokenPipeError): + q.put("item2") + + +def test_sqldata_raw_io(): + mock_sock = MagicMock(spec=SqlDataSocket) + mock_sock.recv_into.return_value = 4 + mock_sock.closed = False + mock_sock.makefile(buffering=0) + + from google.cloud.sql.connector.sqldata_client import SqlDataRawIO + + + raw_io = SqlDataRawIO(mock_sock) + assert raw_io.readable() is True + assert raw_io.writable() is True + assert raw_io.seekable() is False + + buf = bytearray(10) + assert raw_io.readinto(buf) == 4 + mock_sock.recv_into.assert_called_once_with(buf) + + assert raw_io.write(b"data") == 4 + mock_sock.sendall.assert_called_once_with(b"data") + + raw_io.close() + mock_sock.close.assert_called_once() + + +def test_sqldata_socket_direct_sock_delegation(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + ) + mock_direct = MagicMock(spec=socket.socket) + sock._direct_sock = mock_direct + + mock_direct.send.return_value = 4 + mock_direct.recv.return_value = b"resp" + mock_direct.recv_into.return_value = 4 + mock_direct.makefile.return_value = MagicMock() + mock_direct.gettimeout.return_value = 12.0 + mock_direct.getsockopt.return_value = 1 + mock_direct.getsockname.return_value = ("10.0.0.1", 3307) + mock_direct.getpeername.return_value = ("10.0.0.2", 3307) + + sock.sendall(b"test") + mock_direct.sendall.assert_called_once_with(b"test", 0) + + assert sock.send(b"test") == 4 + mock_direct.send.assert_called_once_with(b"test", 0) + + assert sock.recv(1024) == b"resp" + mock_direct.recv.assert_called_once_with(1024, 0) + + buf = bytearray(10) + assert sock.recv_into(buf) == 4 + mock_direct.recv_into.assert_called_once_with(buf, 0, 0) + + + sock.makefile("r") + mock_direct.makefile.assert_called_once() + + sock.settimeout(12.0) + mock_direct.settimeout.assert_called_once_with(12.0) + assert sock.gettimeout() == 12.0 + + sock.setblocking(False) + mock_direct.setblocking.assert_called_once_with(False) + + sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) + mock_direct.setsockopt.assert_called_once_with(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) + + assert sock.getsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY) == 1 + assert sock.getsockname() == ("10.0.0.1", 3307) + assert sock.getpeername() == ("10.0.0.2", 3307) + + sock.shutdown() + assert sock._closed + + +def test_sqldata_socket_makefile_modes(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + ) + + # Unbuffered binary raw mode + f_raw = sock.makefile(mode="rb", buffering=0) + assert hasattr(f_raw, "readinto") + + # Buffered write mode + f_write = sock.makefile(mode="wb") + assert hasattr(f_write, "write") + + # Buffered read/write mode + f_rw = sock.makefile(mode="r+b") + assert hasattr(f_rw, "read") + assert hasattr(f_rw, "write") + + # Text read mode + f_text = sock.makefile(mode="r", encoding="utf-8") + assert hasattr(f_text, "readline") + + # Default read mode without r/w/+ + f_default = sock.makefile(mode="b") + assert hasattr(f_default, "read") + + sock.close() + + +def test_sqldata_socket_edge_cases(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + + on_close_called = False + + def on_close(): + nonlocal on_close_called + on_close_called = True + raise RuntimeError("on_close error") + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + on_close=on_close, + ) + + # sendall with empty bytes + sock.sendall(b"") + + # recv with non-positive bufsize + assert sock.recv(0) == b"" + assert sock.recv(-1) == b"" + + # recv_into with 0 bytes + buf = bytearray(0) + assert sock.recv_into(buf) == 0 + + # Negative timeout + with pytest.raises(ValueError): + sock.settimeout(-1.0) + + # Exception in transport.close and stream.cancel handled cleanly + mock_grpc_client.transport.close.side_effect = RuntimeError("transport close error") + mock_response_stream.cancel.side_effect = RuntimeError("cancel error") + + sock.close() + assert on_close_called is True + # Calling close again is idempotent + sock.close() + + # Recv after closed returns empty bytes + assert sock.recv(10) == b"" + + +def test_sqldata_socket_reader_loop_messages(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + + # Test session_metadata and terminate_session + resp_meta = sqladmin_v1beta4.StreamSqlDataResponse( + session_metadata=sqladmin_v1beta4.SessionMetadata() + ) + resp_term = sqladmin_v1beta4.StreamSqlDataResponse( + terminate_session=sqladmin_v1beta4.TerminateSession() + ) + + def stream_messages(): + yield resp_meta + yield resp_term + + mock_response_stream.__iter__.side_effect = stream_messages + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + on_success=MagicMock(), + ) + + time.sleep(0.1) + assert sock._closed is True + sock.close() + + +def test_sqldata_socket_reader_loop_resource_exhausted(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + + def stream_err(): + time.sleep(0.01) + raise MockRpcError(grpc.StatusCode.RESOURCE_EXHAUSTED) + yield + + mock_response_stream.__iter__.side_effect = stream_err + + resource_exhausted_called = False + + def on_res_exhausted(err): + nonlocal resource_exhausted_called + resource_exhausted_called = True + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + on_resource_exhausted=on_res_exhausted, + ) + + time.sleep(0.1) + assert resource_exhausted_called is True + + # Recv without fallback raises OSError + with pytest.raises(OSError) as exc_info: + sock.recv(10) + assert exc_info.value.errno == 104 # ECONNRESET + sock.close() + + +def test_sqldata_socket_fallback_error_propagation(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + + def stream_err(): + time.sleep(0.01) + raise MockRpcError(grpc.StatusCode.UNAVAILABLE) + yield + + mock_response_stream.__iter__.side_effect = stream_err + + def broken_fallback(): + raise RuntimeError("fallback failed") + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + fallback_fn=broken_fallback, + ) + + with pytest.raises(RuntimeError, match="fallback failed"): + sock.recv(1024) + + sock.close() + + +@pytest.mark.asyncio +async def test_sqldata_client_quota_project_metadata(): + creds = MagicMock(spec=Credentials) + creds.quota_project_id = "cred-quota" + client = SqlDataClient( + endpoint="sqladmin.googleapis.com", + credentials=creds, + quota_project="custom-quota", + ) + + mock_grpc_client = MagicMock() + mock_grpc_client.transport.stream_sql_data = mock_grpc_client.stream_sql_data + mock_stream = MagicMock() + + def stream_gen(): + time.sleep(0.5) + yield sqladmin_v1beta4.StreamSqlDataResponse() + + mock_stream.__iter__.side_effect = stream_gen + mock_grpc_client.stream_sql_data.return_value = mock_stream + + with patch( + "google.cloud.sqladmin_v1beta4.SqlDataServiceClient", + return_value=mock_grpc_client, + ) as mock_client_cls: + await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=MagicMock(), + enable_iam_auth=False, + on_fallback=MagicMock(), + is_fallback_cached=lambda _: False, + ) + + # Check ClientOptions had quota_project_id + client_options = mock_client_cls.call_args[1]["client_options"] + assert client_options.quota_project_id == "custom-quota" + + # Check metadata included x-goog-request-params + metadata = mock_grpc_client.stream_sql_data.call_args[1]["metadata"] + assert any(k == "x-goog-request-params" for k, _ in metadata) + + await client.close() + + +@pytest.mark.asyncio +async def test_sqldata_client_cached_fallback_connect(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="sqladmin.googleapis.com", + credentials=creds, + ) + + mock_conn_info = MagicMock() + mock_conn_info.get_preferred_ips.return_value = ["10.0.0.1"] + mock_ssl_ctx = MagicMock() + mock_conn_info.create_ssl_context = AsyncMock(return_value=mock_ssl_ctx) + get_conn_info = AsyncMock(return_value=mock_conn_info) + + mock_raw_sock = MagicMock(spec=socket.socket) + mock_ssl_sock = MagicMock(spec=socket.socket) + mock_ssl_ctx.wrap_socket.return_value = mock_ssl_sock + + with patch("socket.create_connection", return_value=mock_raw_sock): + sock = await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=MagicMock(), + is_fallback_cached=lambda _: True, + ) + assert sock is mock_ssl_sock + await client.close() + + +@pytest.mark.asyncio +async def test_sqldata_client_direct_socket_factory_no_ips(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="sqladmin.googleapis.com", + credentials=creds, + ) + + from google.cloud.sql.connector.exceptions import CloudSQLIPTypeError + + mock_conn_info = MagicMock() + mock_conn_info.get_preferred_ips.side_effect = CloudSQLIPTypeError("no ips") + get_conn_info = AsyncMock(return_value=mock_conn_info) + + mock_grpc_client = MagicMock() + mock_grpc_client.transport.stream_sql_data = mock_grpc_client.stream_sql_data + mock_grpc_client.stream_sql_data.side_effect = MockRpcError(grpc.StatusCode.UNAVAILABLE) + + with ( + patch( + "google.cloud.sqladmin_v1beta4.SqlDataServiceClient", + return_value=mock_grpc_client, + ), + pytest.raises( + ValueError, + match="Cannot fallback to direct connection: no IP address available.", + ), + ): + await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=MagicMock(), + is_fallback_cached=lambda _: False, + ) + + +@pytest.mark.asyncio +async def test_sqldata_client_direct_socket_factory_ip_retry_and_exhausted(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="sqladmin.googleapis.com", + credentials=creds, + ) + + mock_conn_info = MagicMock() + mock_conn_info.get_preferred_ips.return_value = ["10.0.0.1", "10.0.0.2"] + mock_ssl_ctx = MagicMock() + mock_ssl_sock = MagicMock(spec=socket.socket) + mock_ssl_ctx.wrap_socket.return_value = mock_ssl_sock + mock_conn_info.create_ssl_context = AsyncMock(return_value=mock_ssl_ctx) + get_conn_info = AsyncMock(return_value=mock_conn_info) + + mock_grpc_client = MagicMock() + mock_grpc_client.transport.stream_sql_data = mock_grpc_client.stream_sql_data + mock_grpc_client.stream_sql_data.side_effect = MockRpcError(grpc.StatusCode.UNAVAILABLE) + + # First IP fails, second IP succeeds + with patch( + "google.cloud.sqladmin_v1beta4.SqlDataServiceClient", + return_value=mock_grpc_client, + ), patch( + "socket.create_connection", side_effect=[OSError("conn ref"), MagicMock()] + ): + sock = await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=MagicMock(), + is_fallback_cached=lambda _: False, + ) + assert sock is mock_ssl_sock + + # All IPs fail + with patch( + "google.cloud.sqladmin_v1beta4.SqlDataServiceClient", + return_value=mock_grpc_client, + ), patch( + "socket.create_connection", side_effect=OSError("all ips failed") + ), pytest.raises(OSError, match="all ips failed"): + await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=MagicMock(), + is_fallback_cached=lambda _: False, + ) + + +@pytest.mark.asyncio +async def test_sqldata_client_connect_resource_exhausted(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="sqladmin.googleapis.com", + credentials=creds, + ) + + on_res_exhausted = MagicMock() + mock_conn_info = MagicMock() + mock_conn_info.get_preferred_ips.return_value = ["10.0.0.1"] + mock_conn_info.create_ssl_context = AsyncMock() + get_conn_info = AsyncMock(return_value=mock_conn_info) + + rpc_err = MockRpcError(grpc.StatusCode.RESOURCE_EXHAUSTED) + mock_grpc_client = MagicMock() + mock_grpc_client.transport.stream_sql_data = mock_grpc_client.stream_sql_data + mock_grpc_client.stream_sql_data.side_effect = rpc_err + + with patch( + "google.cloud.sqladmin_v1beta4.SqlDataServiceClient", + return_value=mock_grpc_client, + ): + with pytest.raises(MockRpcError): + await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=MagicMock(), + is_fallback_cached=lambda _: False, + on_resource_exhausted=on_res_exhausted, + ) + + on_res_exhausted.assert_called_once_with(rpc_err) + + +@pytest.mark.asyncio +async def test_sqldata_client_close_exceptions(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="sqladmin.googleapis.com", + credentials=creds, + ) + + mock_sock = MagicMock(spec=SqlDataSocket) + mock_sock.close.side_effect = RuntimeError("sock close error") + client._active_sockets.add(mock_sock) + + broken_cb = MagicMock(side_effect=RuntimeError("cb error")) + client._on_close_callbacks.append(broken_cb) + + # Should not raise exception + await client.close() + mock_sock.close.assert_called_once() + broken_cb.assert_called_once() + + +def test_sqldata_socket_timeout(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + mock_grpc_client.transport.stream_sql_data = mock_grpc_client.stream_sql_data + req_queue = _RequestQueue() + + def stream_blocking(): + time.sleep(2.0) + yield sqladmin_v1beta4.StreamSqlDataResponse() + + mock_response_stream.__iter__.side_effect = stream_blocking + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + timeout=0.05, + ) + + with pytest.raises(socket.timeout): + sock.recv(1024) + + sock.close() + + +@pytest.mark.asyncio +async def test_sqldata_client_connect_success(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="sqladmin.googleapis.com", + credentials=creds, + ) + + mock_grpc_client = MagicMock() + mock_grpc_client.transport.stream_sql_data = mock_grpc_client.stream_sql_data + mock_stream = MagicMock() + + def stream_gen(): + time.sleep(1.0) + yield sqladmin_v1beta4.StreamSqlDataResponse() + + mock_stream.__iter__.side_effect = stream_gen + mock_grpc_client.stream_sql_data.return_value = mock_stream + + with patch( + "google.cloud.sqladmin_v1beta4.SqlDataServiceClient", + return_value=mock_grpc_client, + ): + on_success = MagicMock() + sock = await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=MagicMock(), + enable_iam_auth=False, + on_fallback=MagicMock(), + is_fallback_cached=lambda _: False, + on_success=on_success, + ) + + assert isinstance(sock, SqlDataSocket) + await client.close() + assert sock._closed + + +@pytest.mark.asyncio +async def test_sqldata_client_fallback(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="sqladmin.googleapis.com", + credentials=creds, + ) + + mock_grpc_client = MagicMock() + mock_grpc_client.transport.stream_sql_data = mock_grpc_client.stream_sql_data + rpc_err = MockRpcError(grpc.StatusCode.UNAVAILABLE) + mock_grpc_client.stream_sql_data.side_effect = rpc_err + + mock_conn_info = MagicMock() + mock_conn_info.get_preferred_ips.return_value = ["1.2.3.4"] + mock_ssl_ctx = MagicMock() + mock_conn_info.create_ssl_context = AsyncMock( + return_value=mock_ssl_ctx + ) + get_conn_info = AsyncMock(return_value=mock_conn_info) + + mock_raw_sock = MagicMock(spec=socket.socket) + mock_ssl_sock = MagicMock(spec=socket.socket) + mock_ssl_ctx.wrap_socket.return_value = mock_ssl_sock + + on_fallback = MagicMock() + + with patch( + "google.cloud.sqladmin_v1beta4.SqlDataServiceClient", + return_value=mock_grpc_client, + ), patch( + "socket.create_connection", return_value=mock_raw_sock + ): + sock = await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=on_fallback, + is_fallback_cached=lambda _: False, + ) + + assert sock is mock_ssl_sock + assert on_fallback.called + await client.close() + + +def test_sqldata_socket_transparent_fallback(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + + # Simulate gRPC stream raising FAILED_PRECONDITION on first read + def stream_failing(): + time.sleep(0.05) + raise MockRpcError(grpc.StatusCode.FAILED_PRECONDITION) + yield # Make it a generator + + mock_response_stream.__iter__.side_effect = stream_failing + + mock_direct_sock = MagicMock(spec=socket.socket) + mock_direct_sock.recv.return_value = b"direct server response" + fallback_called = False + + def fallback_fn(): + nonlocal fallback_called + fallback_called = True + return mock_direct_sock + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + timeout=2.0, + fallback_fn=fallback_fn, + ) + + # Client writes startup message before first read + sock.sendall(b"startup message") + + # Client reads response -> triggers fallback, replays write, returns direct response + resp = sock.recv(1024) + + assert fallback_called is True + assert resp == b"direct server response" + mock_direct_sock.sendall.assert_called_once_with(b"startup message") + mock_direct_sock.recv.assert_called_once_with(1024, 0) + sock.close() + + +def test_sqldata_socket_reader_loop_closed(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + + def stream_loop(): + yield sqladmin_v1beta4.StreamSqlDataResponse() + time.sleep(0.05) + yield sqladmin_v1beta4.StreamSqlDataResponse() + + mock_response_stream.__iter__.side_effect = stream_loop + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + timeout=1.0, + ) + sock._closed = True + time.sleep(0.1) + sock.close() + + +def test_sqldata_socket_recv_closed_empty_queue(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + mock_response_stream.__iter__.return_value = iter([]) + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + timeout=0.01, + ) + try: + sock._read_queue.get_nowait() + except queue.Empty: + pass + sock._closed = True + assert sock.recv(1024) == b"" + sock.close() + + +def test_sqldata_socket_close_with_exceptions(): + mock_response_stream = MagicMock() + mock_grpc_client = MagicMock() + req_queue = _RequestQueue() + mock_response_stream.__iter__.return_value = iter([]) + + mock_direct_sock = MagicMock() + mock_direct_sock.close.side_effect = Exception("direct close error") + + sock = SqlDataSocket( + request_queue=req_queue, + response_stream=mock_response_stream, + grpc_client=mock_grpc_client, + timeout=0.01, + ) + sock._direct_sock = mock_direct_sock + + with patch("socket.socket.close", side_effect=Exception("super close error")): + sock.close() + + +@pytest.mark.asyncio +async def test_sqldata_client_fallback_no_ips(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="https://example.com", + credentials=creds, + ) + mock_conn_info = MagicMock() + mock_conn_info.get_preferred_ips.return_value = [] + get_conn_info = AsyncMock(return_value=mock_conn_info) + + with pytest.raises(ValueError, match="no IP address available"): + await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=MagicMock(), + is_fallback_cached=lambda _: True, + ) + + +@pytest.mark.asyncio +async def test_sqldata_client_connect_fallback_fn_invoked(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="https://example.com", + credentials=creds, + ) + + mock_raw_sock = MagicMock(spec=socket.socket) + mock_ssl_sock = MagicMock(spec=socket.socket) + mock_ssl_sock.recv.return_value = b"direct response" + mock_ssl_ctx = MagicMock() + mock_ssl_ctx.wrap_socket.return_value = mock_ssl_sock + + mock_conn_info = MagicMock() + mock_conn_info.get_preferred_ips.return_value = ["1.2.3.4"] + mock_conn_info.create_ssl_context = AsyncMock(return_value=mock_ssl_ctx) + get_conn_info = AsyncMock(return_value=mock_conn_info) + + mock_grpc_client = MagicMock() + del mock_grpc_client.transport + + def stream_failing(*args, **kwargs): + raise MockRpcError(grpc.StatusCode.FAILED_PRECONDITION) + yield + + mock_stream = MagicMock() + mock_stream.__iter__.side_effect = stream_failing + mock_grpc_client.stream_sql_data.return_value = mock_stream + + with patch( + "google.cloud.sqladmin_v1beta4.SqlDataServiceClient", + return_value=mock_grpc_client, + ), patch( + "socket.create_connection", return_value=mock_raw_sock + ): + on_fallback = MagicMock() + sock = await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=on_fallback, + is_fallback_cached=lambda _: False, + ) + resp = sock.recv(1024) + assert resp == b"direct response" + on_fallback.assert_called_once_with("proj:region:inst") + await client.close() + + +@pytest.mark.asyncio +async def test_sqldata_client_connect_fallback_fn_none(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="https://example.com", + credentials=creds, + ) + + mock_conn_info = MagicMock() + mock_conn_info.get_preferred_ips.side_effect = Exception("ip failure") + get_conn_info = AsyncMock(return_value=mock_conn_info) + + mock_grpc_client = MagicMock() + mock_grpc_client.transport.stream_sql_data = mock_grpc_client.stream_sql_data + + def stream_failing(*args, **kwargs): + raise MockRpcError(grpc.StatusCode.FAILED_PRECONDITION) + yield + + mock_stream = MagicMock() + mock_stream.__iter__.side_effect = stream_failing + mock_grpc_client.stream_sql_data.return_value = mock_stream + + with patch( + "google.cloud.sqladmin_v1beta4.SqlDataServiceClient", + return_value=mock_grpc_client, + ): + on_fallback = MagicMock() + sock = await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=on_fallback, + is_fallback_cached=lambda _: False, + ) + with pytest.raises(ValueError, match="No direct fallback connection factory available"): + sock.recv(1024) + await client.close() + + +@pytest.mark.asyncio +async def test_sqldata_client_connect_close_transport_exception(): + creds = MagicMock(spec=Credentials) + client = SqlDataClient( + endpoint="https://example.com", + credentials=creds, + ) + mock_raw_sock = MagicMock(spec=socket.socket) + mock_ssl_sock = MagicMock(spec=socket.socket) + mock_ssl_ctx = MagicMock() + mock_ssl_ctx.wrap_socket.return_value = mock_ssl_sock + + mock_conn_info = MagicMock() + mock_conn_info.get_preferred_ips.return_value = ["1.2.3.4"] + mock_conn_info.create_ssl_context = AsyncMock( + side_effect=[Exception("init prep fail"), mock_ssl_ctx] + ) + get_conn_info = AsyncMock(return_value=mock_conn_info) + + mock_grpc_client = MagicMock() + mock_grpc_client.transport.stream_sql_data.side_effect = Exception("stream fail") + mock_grpc_client.transport.close.side_effect = Exception("close fail") + + with patch( + "google.cloud.sqladmin_v1beta4.SqlDataServiceClient", + return_value=mock_grpc_client, + ), patch( + "socket.create_connection", return_value=mock_raw_sock + ): + sock = await client.connect( + instance_connection_name="proj:region:inst", + region="region", + project="proj", + get_conn_info=get_conn_info, + enable_iam_auth=False, + on_fallback=MagicMock(), + is_fallback_cached=lambda _: False, + ) + assert sock is mock_ssl_sock + await client.close() + + + diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 1c2ae4fd..82322c04 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -75,3 +75,16 @@ def test_format_database_user_mysql() -> None: user2 = utils.format_database_user("MYSQL_8_0", "test") assert user == "test" assert user2 == "test" + + +def test_iptypes_from_str() -> None: + """Test IPTypes._from_str parses string values properly.""" + from google.cloud.sql.connector.enums import IPTypes + + assert IPTypes._from_str("sqldata") == IPTypes.SQL_DATA + assert IPTypes._from_str("sql_data") == IPTypes.SQL_DATA + assert IPTypes._from_str("public") == IPTypes.PUBLIC + assert IPTypes._from_str("private") == IPTypes.PRIVATE + assert IPTypes._from_str("psc") == IPTypes.PSC + +