diff --git a/asyncpg/connect_utils.py b/asyncpg/connect_utils.py index 07c4fdde..0db65826 100644 --- a/asyncpg/connect_utils.py +++ b/asyncpg/connect_utils.py @@ -613,6 +613,14 @@ def _parse_connect_dsn_and_args(*, dsn, host, port, user, raise exceptions.ClientConfigurationError( 'could not determine database name to connect to') + # The startup packet is built by pgproto's WriteBuffer.write_str(), + # which is typed to accept exactly `str` and does not accept `str` + # subclasses (e.g. enum.StrEnum members, or third-party string-like + # types such as tomlkit's), even though `isinstance(x, str)` is True + # for them. Coerce here so any such value is safely accepted. + user = str(user) + database = str(database) + if password is None: if passfile is None: passfile = os.getenv('PGPASSFILE') @@ -821,6 +829,11 @@ def _parse_connect_dsn_and_args(*, dsn, host, port, user, raise exceptions.ClientConfigurationError( 'server_settings is expected to be None or ' 'a Dict[str, str]') + if server_settings is not None: + # See the comment above the `user`/`database` coercion: keys and + # values are also written via pgproto's write_str(), which rejects + # str subclasses. + server_settings = {str(k): str(v) for k, v in server_settings.items()} if target_session_attrs is None: target_session_attrs = os.getenv( diff --git a/tests/test_connect.py b/tests/test_connect.py index 955fb825..edfe54c3 100644 --- a/tests/test_connect.py +++ b/tests/test_connect.py @@ -1256,6 +1256,46 @@ def test_connect_params(self): for testcase in self.TESTS: self.run_testcase(testcase) + def test_connect_params_coerces_str_subclasses(self): + # user/database/server_settings are later written by pgproto's + # WriteBuffer.write_str(), which is typed to accept exactly `str` + # and rejects str subclasses (e.g. enum.StrEnum members) even + # though isinstance(x, str) is True for them. _parse_connect_dsn_ + # and_args must coerce these to plain str so such values are + # safely accepted instead of blowing up deep in the protocol + # layer. See #1340. + class SUser(str): + pass + + class SDb(str): + pass + + class SKey(str): + pass + + class SVal(str): + pass + + user = SUser('someuser') + database = SDb('somedb') + server_settings = {SKey('application_name'): SVal('someapp')} + + self.assertIsInstance(user, str) + self.assertNotEqual(type(user), str) + + _, params = connect_utils._parse_connect_dsn_and_args( + dsn=None, host=None, port=None, user=user, password=None, + passfile=None, database=database, ssl=None, + direct_tls=False, server_settings=server_settings, + target_session_attrs=None, krbsrvname=None, gsslib=None, + service=None, servicefile=None) + + self.assertEqual(type(params.user), str) + self.assertEqual(type(params.database), str) + for k, v in params.server_settings.items(): + self.assertEqual(type(k), str) + self.assertEqual(type(v), str) + def test_connect_connection_service_file(self): connection_service_file = tempfile.NamedTemporaryFile( 'w+t', delete=False)