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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion google/cloud/sql/connector/asyncpg.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,9 @@ async def connect(
'Unable to import module "asyncpg." Please install and try again.'
)
user = kwargs.pop("user")
db = kwargs.pop("db")
db = kwargs.pop("database", kwargs.pop("db", None))
if db is None:
raise KeyError("database")
passwd = kwargs.pop("password", None)

return await asyncpg.connect(
Expand Down
4 changes: 3 additions & 1 deletion google/cloud/sql/connector/pg8000.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,9 @@ def connect(
)

user = kwargs.pop("user")
db = kwargs.pop("db")
db = kwargs.pop("database", kwargs.pop("db", None))
if db is None:
raise KeyError("database")
passwd = kwargs.pop("password", None)
return pg8000.dbapi.connect(
user,
Expand Down
6 changes: 6 additions & 0 deletions google/cloud/sql/connector/pymysql.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,12 @@ def connect(
# pop timeout as timeout arg is called 'connect_timeout' for pymysql
timeout = kwargs.pop("timeout")
kwargs["connect_timeout"] = kwargs.get("connect_timeout", timeout)

# map 'db' to 'database' to avoid deprecation warning in pymysql
db = kwargs.pop("db", None)
if db is not None:
kwargs["database"] = db

# Create pymysql connection object and hand in pre-made connection
conn = pymysql.Connection(host=ip_address, defer_connect=True, **kwargs)
conn.connect(sock)
Expand Down
2 changes: 1 addition & 1 deletion google/cloud/sql/connector/pytds.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,7 @@ def connect(ip_address: str, sock: ssl.SSLSocket, **kwargs: Any) -> "pytds.Conne
'Unable to import module "pytds." Please install and try again.'
)

db = kwargs.pop("db", None)
db = kwargs.pop("database", kwargs.pop("db", None))

if kwargs.pop("active_directory_auth", False):
if platform.system() == "Windows":
Expand Down
4 changes: 2 additions & 2 deletions tests/system/test_asyncpg_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,7 @@ async def create_sqlalchemy_engine(
"asyncpg",
user=user,
password=password,
db=db,
database=db,
ip_type=ip_type, # can be "public", "private" or "psc"
**kwargs, # additional asyncpg connection args
),
Expand Down Expand Up @@ -155,7 +155,7 @@ async def getconn(
"asyncpg",
user=user,
password=password,
db=db,
database=db,
ip_type=ip_type, # can be "public", "private" or "psc"
**kwargs,
)
Expand Down
2 changes: 1 addition & 1 deletion tests/system/test_asyncpg_iam_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,7 @@ async def create_sqlalchemy_engine(
instance_connection_name,
"asyncpg",
user=user,
db=db,
database=db,
ip_type=ip_type, # can be "public", "private" or "psc"
enable_iam_auth=True,
),
Expand Down
4 changes: 2 additions & 2 deletions tests/system/test_connector_object.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ def getconn() -> pymysql.connections.Connection:
"pymysql",
user=os.environ["MYSQL_USER"],
password=os.environ["MYSQL_PASS"],
db=os.environ["MYSQL_DB"],
database=os.environ["MYSQL_DB"],
ip_type=os.environ.get("IP_TYPE", "public"),
)
return conn
Expand Down Expand Up @@ -142,5 +142,5 @@ def test_connector_sqlserver_iam_auth_error() -> None:
"pytds",
user="my-user",
password="my-pass",
db="my-db",
database="my-db",
)
2 changes: 1 addition & 1 deletion tests/system/test_ip_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ def getconn() -> pymysql.connections.Connection:
ip_type=ip_type,
user=os.environ["MYSQL_USER"],
password=os.environ["MYSQL_PASS"],
db=os.environ["MYSQL_DB"],
database=os.environ["MYSQL_DB"],
)
return conn

Expand Down
2 changes: 1 addition & 1 deletion tests/system/test_pg8000_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,7 +87,7 @@ def create_sqlalchemy_engine(
"pg8000",
user=user,
password=password,
db=db,
database=db,
ip_type=ip_type, # can be "public", "private" or "psc"
),
)
Expand Down
2 changes: 1 addition & 1 deletion tests/system/test_pg8000_iam_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ def create_sqlalchemy_engine(
instance_connection_name,
"pg8000",
user=user,
db=db,
database=db,
ip_type=ip_type, # can be "public", "private" or "psc"
enable_iam_auth=True,
),
Expand Down
2 changes: 1 addition & 1 deletion tests/system/test_pymysql_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,7 +78,7 @@ def create_sqlalchemy_engine(
"pymysql",
user=user,
password=password,
db=db,
database=db,
ip_type=ip_type, # can be "public", "private" or "psc"
),
)
Expand Down
2 changes: 1 addition & 1 deletion tests/system/test_pymysql_iam_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ def create_sqlalchemy_engine(
instance_connection_name,
"pymysql",
user=user,
db=db,
database=db,
ip_type=ip_type, # can be "public", "private" or "psc"
enable_iam_auth=True,
),
Expand Down
2 changes: 1 addition & 1 deletion tests/system/test_pytds_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,7 +76,7 @@ def create_sqlalchemy_engine(
"pytds",
user=user,
password=password,
db=db,
database=db,
ip_type=ip_type, # can be "public", "private" or "psc"
),
)
Expand Down
61 changes: 61 additions & 0 deletions tests/unit/test_asyncpg.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,3 +35,64 @@ async def test_asyncpg(mock_connect: AsyncMock, kwargs: Any) -> None:
assert connection is True
# verify that driver connection call would be made
assert mock_connect.assert_called_once


@pytest.mark.asyncio
async def test_asyncpg_import_error(kwargs: Any) -> None:
"""Test to verify that connect raises ImportError if asyncpg is not installed."""
with patch.dict("sys.modules", {"asyncpg": None}):
with pytest.raises(ImportError) as excinfo:
await connect("0.0.0.0", ssl.create_default_context(), **kwargs)
assert 'Unable to import module "asyncpg."' in str(excinfo.value)


@pytest.mark.asyncio
async def test_asyncpg_database_param(kwargs: Any) -> None:
"""Test that asyncpg wrapper accepts both 'db' and 'database'."""
ip_addr = "0.0.0.0"
context = ssl.create_default_context()

# Test with 'db'
kwargs_db = kwargs.copy()
kwargs_db["db"] = "my-db-1"
with patch("asyncpg.connect", new_callable=AsyncMock) as mock_connect:
await connect(ip_addr, context, **kwargs_db)
mock_connect.assert_called_once_with(
user=kwargs["user"],
database="my-db-1",
password=kwargs.get("password"),
host=ip_addr,
port=3307,
ssl=context,
direct_tls=True,
)

# Test with 'database'
kwargs_database = kwargs.copy()
kwargs_database["database"] = "my-db-2"
if "db" in kwargs_database:
del kwargs_database["db"]
with patch("asyncpg.connect", new_callable=AsyncMock) as mock_connect:
await connect(ip_addr, context, **kwargs_database)
mock_connect.assert_called_once_with(
user=kwargs["user"],
database="my-db-2",
password=kwargs.get("password"),
host=ip_addr,
port=3307,
ssl=context,
direct_tls=True,
)


@pytest.mark.asyncio
async def test_asyncpg_database_missing(kwargs: Any) -> None:
"""Test that asyncpg wrapper raises KeyError if database is missing."""
if "db" in kwargs:
del kwargs["db"]
if "database" in kwargs:
del kwargs["database"]
with pytest.raises(KeyError, match="database"):
await connect("0.0.0.0", ssl.create_default_context(), **kwargs)


Loading
Loading