Skip to content

Commit f1ca84b

Browse files
authored
[DBMON-5768] Avoid collision between multiple TokenProviders (DataDog#21560)
* Avoid collision between multiple TokenProviders * Add changelog
1 parent a85074c commit f1ca84b

4 files changed

Lines changed: 130 additions & 9 deletions

File tree

postgres/changelog.d/21560.fixed

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
Fixes a collision issue when token based authentication is configured for multiple Postgres instances

postgres/datadog_checks/postgres/connection_pool.py

Lines changed: 11 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -97,15 +97,17 @@ class TokenAwareConnection(Connection):
9797
Connection that can be used for managed authentication.
9898
"""
9999

100-
token_provider: Optional[TokenProvider] = None
101-
102100
@classmethod
103101
def connect(cls, *args, **kwargs):
104102
"""
105103
Override the connection method to pass a refreshable token as the connection password.
104+
105+
The token_provider can be passed via the 'token_provider' kwarg and will be used
106+
to dynamically fetch authentication tokens.
106107
"""
107-
if cls.token_provider:
108-
kwargs["password"] = cls.token_provider.get_token()
108+
token_provider = kwargs.pop("token_provider", None)
109+
if token_provider:
110+
kwargs["password"] = token_provider.get_token()
109111
return super().connect(*args, **kwargs)
110112

111113

@@ -193,6 +195,7 @@ def __init__(
193195
pool_config (dict, optional): Additional ConnectionPool settings (min_size, max_size, etc).
194196
statement_timeout (int, optional): Statement timeout in milliseconds.
195197
sqlascii_encodings (list[str], optional): List of encodings to handle for SQLASCII text.
198+
token_provider (TokenProvider, optional): Token provider for managed authentication.
196199
"""
197200
self.max_db = max_db
198201
self.base_conn_args = base_conn_args
@@ -207,8 +210,6 @@ def __init__(
207210
"open": True,
208211
}
209212

210-
TokenAwareConnection.token_provider = self.token_provider
211-
212213
self.lock = threading.Lock()
213214
self.pools: OrderedDict[str, Tuple[ConnectionPool, float, bool]] = OrderedDict()
214215
self._closed = False
@@ -240,6 +241,10 @@ def _create_pool(self, dbname: str) -> ConnectionPool:
240241
"""
241242
kwargs = self.base_conn_args.as_kwargs(dbname=dbname)
242243

244+
# Pass the token_provider as a kwarg so it's available to TokenAwareConnection.connect()
245+
if self.token_provider:
246+
kwargs["token_provider"] = self.token_provider
247+
243248
return ConnectionPool(
244249
kwargs=kwargs,
245250
configure=self._configure_connection,

postgres/datadog_checks/postgres/postgres.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -959,7 +959,13 @@ def _new_connection(self, dbname):
959959
# TODO: Keeping this main connection outside of the pool for now to keep existing behavior.
960960
# We should move this to the pool in the future.
961961
conn_args = self.build_connection_args()
962-
conn = TokenAwareConnection.connect(**conn_args.as_kwargs(dbname=dbname))
962+
kwargs = conn_args.as_kwargs(dbname=dbname)
963+
964+
# Pass the token_provider as a kwarg so it's available to TokenAwareConnection.connect()
965+
if self.db_pool.token_provider:
966+
kwargs["token_provider"] = self.db_pool.token_provider
967+
968+
conn = TokenAwareConnection.connect(**kwargs)
963969
self.db_pool._configure_connection(conn)
964970
return conn
965971

postgres/tests/test_token_provider.py

Lines changed: 111 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,9 +3,15 @@
33
# Licensed under a 3-clause BSD style license (see LICENSE)
44

55
import time
6-
from unittest.mock import Mock, patch
6+
from unittest.mock import MagicMock, Mock, patch
77

8-
from datadog_checks.postgres.connection_pool import AWSTokenProvider, AzureTokenProvider, TokenProvider
8+
from datadog_checks.postgres.connection_pool import (
9+
AWSTokenProvider,
10+
AzureTokenProvider,
11+
LRUConnectionPoolManager,
12+
PostgresConnectionArgs,
13+
TokenProvider,
14+
)
915

1016

1117
def test_get_token_first_call():
@@ -256,6 +262,109 @@ def test_azure_token_provider_integration():
256262
assert mock_credential.get_token.call_count == 1
257263

258264

265+
def test_multiple_connection_pools_no_token_collision():
266+
"""
267+
Test that multiple LRUConnectionPoolManager instances with different token providers
268+
don't have token collision issues.
269+
"""
270+
# Create two different mock token providers that return different tokens
271+
token_provider_1 = MockTokenProvider()
272+
token_provider_1._fetch_token = Mock(return_value=("token_for_rds_1", time.time() + 3600))
273+
274+
token_provider_2 = MockTokenProvider()
275+
token_provider_2._fetch_token = Mock(return_value=("token_for_rds_2", time.time() + 3600))
276+
277+
# Create connection args for two different RDS instances
278+
conn_args_1 = PostgresConnectionArgs(
279+
application_name="test_rds_1",
280+
username="user1",
281+
host="rds-instance-1.amazonaws.com",
282+
port=5432,
283+
)
284+
285+
conn_args_2 = PostgresConnectionArgs(
286+
application_name="test_rds_2",
287+
username="user2",
288+
host="rds-instance-2.amazonaws.com",
289+
port=5432,
290+
)
291+
292+
# Create two connection pool managers with different token providers
293+
pool_manager_1 = LRUConnectionPoolManager(
294+
max_db=2,
295+
base_conn_args=conn_args_1,
296+
token_provider=token_provider_1,
297+
)
298+
299+
pool_manager_2 = LRUConnectionPoolManager(
300+
max_db=2,
301+
base_conn_args=conn_args_2,
302+
token_provider=token_provider_2,
303+
)
304+
305+
# Mock the ConnectionPool to capture what connection_class is passed
306+
# and simulate connections to verify the token provider behavior
307+
captured_pools = []
308+
309+
def mock_ConnectionPool(*args, **pool_kwargs):
310+
# Capture the kwargs and connection_class
311+
connection_class = pool_kwargs.get('connection_class')
312+
conn_kwargs = pool_kwargs.get('kwargs', {})
313+
314+
captured_pools.append(
315+
{
316+
'connection_class': connection_class,
317+
'kwargs': conn_kwargs.copy() if conn_kwargs else {},
318+
}
319+
)
320+
321+
# Create a mock pool
322+
mock_pool = MagicMock()
323+
return mock_pool
324+
325+
with patch('datadog_checks.postgres.connection_pool.ConnectionPool', side_effect=mock_ConnectionPool):
326+
# Create pools
327+
pool_manager_1._create_pool("db1")
328+
pool_manager_2._create_pool("db2")
329+
330+
# Verify that pools were created
331+
assert len(captured_pools) == 2, "Should have created 2 connection pools"
332+
333+
# Get the pools for each RDS instance
334+
pool_1 = captured_pools[0]
335+
pool_2 = captured_pools[1]
336+
337+
# Simulate what happens when a connection is established by calling the connect method
338+
# This will trigger the token provider logic
339+
conn_class_1 = pool_1['connection_class']
340+
conn_class_2 = pool_2['connection_class']
341+
342+
# Call connect to see what token each would use
343+
# We need to mock the parent connect to capture the password
344+
with patch('psycopg.Connection.connect') as mock_parent_connect:
345+
mock_parent_connect.return_value = MagicMock()
346+
347+
# Call connect for pool 1 with its kwargs
348+
conn_class_1.connect(**pool_1['kwargs'])
349+
password_1 = mock_parent_connect.call_args[1].get('password')
350+
351+
# Call connect for pool 2 with its kwargs
352+
conn_class_2.connect(**pool_2['kwargs'])
353+
password_2 = mock_parent_connect.call_args[1].get('password')
354+
355+
# Verify no collision - each pool manager uses its own token provider
356+
assert password_1 == "token_for_rds_1"
357+
assert password_2 == "token_for_rds_2"
358+
359+
# Verify each token provider was called independently
360+
assert token_provider_1._fetch_token.call_count >= 1, "Token provider 1 should be called at least once"
361+
assert token_provider_2._fetch_token.call_count >= 1, "Token provider 2 should be called at least once"
362+
363+
# Clean up
364+
pool_manager_1.close_all()
365+
pool_manager_2.close_all()
366+
367+
259368
class MockTokenProvider(TokenProvider):
260369
"""Mock implementation of TokenProvider for testing."""
261370

0 commit comments

Comments
 (0)