Skip to content

Commit 16e2e53

Browse files
committed
feat(sqldata): prefer PRIVATE, PSC, PUBLIC IP order on direct fallback
1 parent 6de01f5 commit 16e2e53

2 files changed

Lines changed: 71 additions & 7 deletions

File tree

google/cloud/sql/connector/sqldata_client.py

Lines changed: 21 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -142,6 +142,20 @@ async def close(self) -> None:
142142
except Exception: # noqa: BLE001, S110
143143
pass
144144

145+
async def _open_direct_connection(
146+
self,
147+
target_ip: str,
148+
port: int,
149+
ssl_context: Any,
150+
connect_timeout: float,
151+
) -> tuple[asyncio.StreamReader, asyncio.StreamWriter]:
152+
return await asyncio.wait_for(
153+
asyncio.open_connection(
154+
target_ip, port, ssl=ssl_context, server_hostname=target_ip
155+
),
156+
timeout=connect_timeout,
157+
)
158+
145159
async def _handle_tunnel(
146160
self,
147161
client_reader: asyncio.StreamReader,
@@ -225,9 +239,9 @@ async def connect_grpc() -> tuple[grpc.aio.Channel, Any]:
225239
async def connect_direct() -> tuple[asyncio.StreamReader, asyncio.StreamWriter]:
226240
logger.debug("Fallback triggered, fetching connection info...")
227241
conn_info = await get_conn_info()
228-
# Find a fallback IP address, prioritizing PUBLIC for direct fallback connectivity
242+
# Find a fallback IP address, prioritizing PRIVATE, PSC, PUBLIC
229243
targets: list[str] = []
230-
for t in [IPTypes.PUBLIC, IPTypes.PSC, IPTypes.PRIVATE]:
244+
for t in [IPTypes.PRIVATE, IPTypes.PSC, IPTypes.PUBLIC]:
231245
try:
232246
targets.extend(conn_info.get_preferred_ips(t))
233247
except CloudSQLIPTypeError as e:
@@ -240,11 +254,11 @@ async def connect_direct() -> tuple[asyncio.StreamReader, asyncio.StreamWriter]:
240254
for target_ip in targets:
241255
logger.debug(f"Connecting directly to {target_ip}:{SERVER_PROXY_PORT}")
242256
try:
243-
r, w = await asyncio.wait_for(
244-
asyncio.open_connection(
245-
target_ip, SERVER_PROXY_PORT, ssl=ssl_context, server_hostname=target_ip
246-
),
247-
timeout=connect_timeout,
257+
r, w = await self._open_direct_connection(
258+
target_ip,
259+
SERVER_PROXY_PORT,
260+
ssl_context,
261+
connect_timeout,
248262
)
249263
self._active_writers.add(w)
250264
return r, w

tests/unit/test_connector.py

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1080,5 +1080,55 @@ async def mock_connect_tunnel(**kwargs):
10801080
assert state.last_err is None
10811081

10821082

1083+
@pytest.mark.asyncio
1084+
async def test_sqldata_fallback_ip_order(fake_credentials: Credentials) -> None:
1085+
"""Test that direct fallback queries IP addresses in PRIVATE, PSC, PUBLIC order."""
1086+
client = SqlDataClient(
1087+
endpoint="sqladmin.googleapis.com",
1088+
credentials=fake_credentials,
1089+
)
1090+
mock_conn_info = MagicMock()
1091+
queried_ip_types: list[IPTypes] = []
1092+
1093+
def mock_get_preferred_ips(ip_type: IPTypes):
1094+
queried_ip_types.append(ip_type)
1095+
if ip_type == IPTypes.PUBLIC:
1096+
return ["1.2.3.4"]
1097+
from google.cloud.sql.connector.exceptions import CloudSQLIPTypeError
1098+
1099+
raise CloudSQLIPTypeError(f"{ip_type} not available")
1100+
1101+
mock_conn_info.get_preferred_ips.side_effect = mock_get_preferred_ips
1102+
mock_conn_info.create_ssl_context = AsyncMock(return_value=None)
1103+
get_conn_info = AsyncMock(return_value=mock_conn_info)
1104+
1105+
mock_reader = AsyncMock()
1106+
mock_reader.read = AsyncMock(return_value=b"")
1107+
mock_writer = MagicMock()
1108+
mock_writer.wait_closed = AsyncMock()
1109+
client._open_direct_connection = AsyncMock(
1110+
return_value=(mock_reader, mock_writer)
1111+
)
1112+
1113+
port = await client.connect_tunnel(
1114+
instance_connection_name="proj:reg:inst",
1115+
region="reg",
1116+
project="proj",
1117+
get_conn_info=get_conn_info,
1118+
enable_iam_auth=False,
1119+
on_fallback=MagicMock(),
1120+
is_fallback_cached=MagicMock(return_value=True),
1121+
)
1122+
1123+
# Trigger client connection to tunnel
1124+
_r, w = await asyncio.open_connection("127.0.0.1", port)
1125+
await asyncio.sleep(0.1)
1126+
w.close()
1127+
await w.wait_closed()
1128+
1129+
assert queried_ip_types == [IPTypes.PRIVATE, IPTypes.PSC, IPTypes.PUBLIC]
1130+
await client.close()
1131+
1132+
10831133

10841134

0 commit comments

Comments
 (0)