diff --git a/gql/transport/requests.py b/gql/transport/requests.py index 4e4e6ffb..996e6909 100644 --- a/gql/transport/requests.py +++ b/gql/transport/requests.py @@ -123,6 +123,9 @@ def connect(self): # Creating a session that can later be re-use to configure custom mechanisms self.session = requests.Session() + if self.headers: + self.session.headers.update(self.headers) + # If we specified some retries, we provide a predefined retry-logic if self.retries > 0: adapter = HTTPAdapter( diff --git a/tests/test_aiohttp.py b/tests/test_aiohttp.py index 00bd8a0f..7ba9b3ec 100644 --- a/tests/test_aiohttp.py +++ b/tests/test_aiohttp.py @@ -1918,3 +1918,25 @@ async def handler(request): await session.execute("qmlsdkfj") assert "request should be a GraphQLRequest object" in str(exc_info.value) + + +@pytest.mark.asyncio +async def test_aiohttp_save_headers_in_session(): + """Regression test for issue #613""" + from gql.transport.aiohttp import AIOHTTPTransport + + transport = AIOHTTPTransport("url", headers={"test": "header"}) + await transport.connect() + assert transport.session + assert transport.session.headers["test"] == "header" + + transport2 = AIOHTTPTransport("url") + await transport2.connect() + assert transport2.session + + del transport.session.headers["test"] + + assert transport.session.headers == transport2.session.headers + + await transport.close() + await transport2.close() diff --git a/tests/test_requests.py b/tests/test_requests.py index 7de4a12a..4cb8fb10 100644 --- a/tests/test_requests.py +++ b/tests/test_requests.py @@ -1272,3 +1272,24 @@ def test_code(): assert pi == Decimal("3.141592653589793238462643383279502884197") await run_sync_test(server, test_code) + + +def test_requests_save_headers_in_session(): + """Regression test for issue #613""" + from gql.transport.requests import RequestsHTTPTransport + + transport = RequestsHTTPTransport("url", headers={"test": "header"}) + transport.connect() + assert transport.session + assert transport.session.headers["test"] == "header" + + transport2 = RequestsHTTPTransport("url") + transport2.connect() + assert transport2.session + + del transport.session.headers["test"] + + assert transport.session.headers == transport2.session.headers + + transport.close() + transport2.close()