diff --git a/httpcore/client.py b/httpcore/client.py index cb8ead9fbd..d4dcd33f16 100644 --- a/httpcore/client.py +++ b/httpcore/client.py @@ -31,6 +31,7 @@ class AsyncClient: def __init__( self, + base_url: URLTypes = None, ssl: SSLConfig = DEFAULT_SSL_CONFIG, timeout: TimeoutConfig = DEFAULT_TIMEOUT_CONFIG, pool_limits: PoolLimits = DEFAULT_POOL_LIMITS, @@ -46,6 +47,8 @@ def __init__( self.max_redirects = max_redirects self.dispatch = dispatch + self.base_url = None if base_url is None else URL(base_url) + async def get( self, url: URLTypes, @@ -221,6 +224,10 @@ async def request( ssl: SSLConfig = None, timeout: TimeoutConfig = None, ) -> Response: + + if self.base_url is not None: + url = URL(url, allow_relative=True).resolve_with(self.base_url) + request = Request( method, url, data=data, query_params=query_params, headers=headers ) @@ -375,6 +382,7 @@ async def __aexit__( class Client: def __init__( self, + base_url: URLTypes = None, ssl: SSLConfig = DEFAULT_SSL_CONFIG, timeout: TimeoutConfig = DEFAULT_TIMEOUT_CONFIG, pool_limits: PoolLimits = DEFAULT_POOL_LIMITS, @@ -382,8 +390,11 @@ def __init__( dispatch: Dispatcher = None, backend: ConcurrencyBackend = None, ) -> None: + self.base_url = None if base_url is None else URL(base_url) + self._client = AsyncClient( ssl=ssl, + base_url=base_url, timeout=timeout, pool_limits=pool_limits, max_redirects=max_redirects, @@ -405,6 +416,10 @@ def request( ssl: SSLConfig = None, timeout: TimeoutConfig = None, ) -> SyncResponse: + + if self.base_url is not None: + url = URL(url, allow_relative=True).resolve_with(self.base_url) + request = Request( method, url, data=data, query_params=query_params, headers=headers ) diff --git a/tests/client/test_async_client.py b/tests/client/test_async_client.py index 27602d12d6..6755258c92 100644 --- a/tests/client/test_async_client.py +++ b/tests/client/test_async_client.py @@ -15,6 +15,16 @@ async def test_get(server): assert repr(response) == "" +@pytest.mark.asyncio +async def test_get_base_url(server): + base_url = "http://127.0.0.1:8000/" + async with httpcore.AsyncClient(base_url=base_url) as client: + response = await client.get("/hello") + assert response.status_code == 200 + assert response.text == "Hello, world!" + assert str(response.url) == "http://127.0.0.1:8000/hello" + + @pytest.mark.asyncio async def test_post(server): url = "http://127.0.0.1:8000/" diff --git a/tests/client/test_client.py b/tests/client/test_client.py index 820efec175..d49833701c 100644 --- a/tests/client/test_client.py +++ b/tests/client/test_client.py @@ -40,6 +40,16 @@ def test_get(server): assert repr(response) == "" +@threadpool +def test_get_base_url(server): + base_url = "http://127.0.0.1:8000/" + with httpcore.Client(base_url=base_url) as client: + response = client.get("/hello") + assert response.status_code == 200 + assert response.text == "Hello, world!" + assert str(response.url) == "http://127.0.0.1:8000/hello" + + @threadpool def test_post(server): with httpcore.Client() as http: