Skip to content

Commit 648c67d

Browse files
committed
Coerce user/database/server_settings to plain str before the wire
WriteBuffer.write_str() (in the vendored pgproto submodule) is typed to accept exactly str and does not accept str subclasses -- e.g. enum.StrEnum members, or third-party string-like values such as tomlkit's -- even though isinstance(x, str) is True for them. When one of these reaches the startup packet via user/database/server_settings, building it raises a TypeError deep inside the compiled protocol layer, which in turn triggers a secondary AttributeError ('Protocol' object has no attribute '_on_error') while trying to report the original failure, masking the real cause entirely. Coerce user, database, and server_settings keys/values to plain str right after they're resolved/validated in _parse_connect_dsn_and_args(), before they ever reach the protocol layer. Verified against a real PostgreSQL server (Docker) with an enum.StrEnum user and server_settings entry: TypeError/AttributeError before the fix, successful connection after. Fixes #1340.
1 parent db8ecc2 commit 648c67d

2 files changed

Lines changed: 53 additions & 0 deletions

File tree

asyncpg/connect_utils.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -613,6 +613,14 @@ def _parse_connect_dsn_and_args(*, dsn, host, port, user,
613613
raise exceptions.ClientConfigurationError(
614614
'could not determine database name to connect to')
615615

616+
# The startup packet is built by pgproto's WriteBuffer.write_str(),
617+
# which is typed to accept exactly `str` and does not accept `str`
618+
# subclasses (e.g. enum.StrEnum members, or third-party string-like
619+
# types such as tomlkit's), even though `isinstance(x, str)` is True
620+
# for them. Coerce here so any such value is safely accepted.
621+
user = str(user)
622+
database = str(database)
623+
616624
if password is None:
617625
if passfile is None:
618626
passfile = os.getenv('PGPASSFILE')
@@ -821,6 +829,11 @@ def _parse_connect_dsn_and_args(*, dsn, host, port, user,
821829
raise exceptions.ClientConfigurationError(
822830
'server_settings is expected to be None or '
823831
'a Dict[str, str]')
832+
if server_settings is not None:
833+
# See the comment above the `user`/`database` coercion: keys and
834+
# values are also written via pgproto's write_str(), which rejects
835+
# str subclasses.
836+
server_settings = {str(k): str(v) for k, v in server_settings.items()}
824837

825838
if target_session_attrs is None:
826839
target_session_attrs = os.getenv(

tests/test_connect.py

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1256,6 +1256,46 @@ def test_connect_params(self):
12561256
for testcase in self.TESTS:
12571257
self.run_testcase(testcase)
12581258

1259+
def test_connect_params_coerces_str_subclasses(self):
1260+
# user/database/server_settings are later written by pgproto's
1261+
# WriteBuffer.write_str(), which is typed to accept exactly `str`
1262+
# and rejects str subclasses (e.g. enum.StrEnum members) even
1263+
# though isinstance(x, str) is True for them. _parse_connect_dsn_
1264+
# and_args must coerce these to plain str so such values are
1265+
# safely accepted instead of blowing up deep in the protocol
1266+
# layer. See #1340.
1267+
class SUser(str):
1268+
pass
1269+
1270+
class SDb(str):
1271+
pass
1272+
1273+
class SKey(str):
1274+
pass
1275+
1276+
class SVal(str):
1277+
pass
1278+
1279+
user = SUser('someuser')
1280+
database = SDb('somedb')
1281+
server_settings = {SKey('application_name'): SVal('someapp')}
1282+
1283+
self.assertIsInstance(user, str)
1284+
self.assertNotEqual(type(user), str)
1285+
1286+
_, params = connect_utils._parse_connect_dsn_and_args(
1287+
dsn=None, host=None, port=None, user=user, password=None,
1288+
passfile=None, database=database, ssl=None,
1289+
direct_tls=False, server_settings=server_settings,
1290+
target_session_attrs=None, krbsrvname=None, gsslib=None,
1291+
service=None, servicefile=None)
1292+
1293+
self.assertEqual(type(params.user), str)
1294+
self.assertEqual(type(params.database), str)
1295+
for k, v in params.server_settings.items():
1296+
self.assertEqual(type(k), str)
1297+
self.assertEqual(type(v), str)
1298+
12591299
def test_connect_connection_service_file(self):
12601300
connection_service_file = tempfile.NamedTemporaryFile(
12611301
'w+t', delete=False)

0 commit comments

Comments
 (0)