diff --git a/poetry.lock b/poetry.lock index f71e9a7c4..e29611082 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,4 +1,4 @@ -# This file is automatically @generated by Poetry 2.4.1 and should not be changed by hand. +# This file is automatically @generated by Poetry 2.2.1 and should not be changed by hand. [[package]] name = "aiohappyeyeballs" @@ -1112,6 +1112,29 @@ http2 = ["h2 (>=3,<5)"] socks = ["socksio (==1.*)"] trio = ["trio (>=0.22.0,<1.0)"] +[[package]] +name = "httpcore2" +version = "2.12.0" +description = "A minimal low-level HTTP client." +optional = false +python-versions = ">=3.10" +groups = ["main"] +markers = "sys_platform != \"emscripten\"" +files = [ + {file = "httpcore2-2.12.0-py3-none-any.whl", hash = "sha256:7e04258ce01013d7d615e5b910a3b27fac937d7a95038227e79652b4ba3b4ceb"}, + {file = "httpcore2-2.12.0.tar.gz", hash = "sha256:9293522bba0aa7c4c8e9e3f040c16575bd8868e155a77fa30c7a9085a5eae648"}, +] + +[package.dependencies] +h11 = ">=0.16" +truststore = ">=0.10" + +[package.extras] +asyncio = ["anyio (>=4.5.0,<5.0)"] +http2 = ["h2 (>=3,<5)"] +socks = ["socksio (==1.*)"] +trio = ["trio (>=0.33.0,<1.0)"] + [[package]] name = "httpx" version = "0.28.1" @@ -1154,6 +1177,47 @@ files = [ aiohttp = ">=3.10.0,<4" httpx = ">=0.27.0" +[[package]] +name = "httpx2" +version = "2.12.0" +description = "The next generation HTTP client." +optional = false +python-versions = ">=3.10" +groups = ["main"] +files = [ + {file = "httpx2-2.12.0-py3-none-any.whl", hash = "sha256:cc8b6eecb8661c146b8f89a60e97456ee086e91a784ed31ac450c3a9e613dd36"}, + {file = "httpx2-2.12.0.tar.gz", hash = "sha256:7631fe9887a8a2275f4a2540e053aa670fcc50742864a9ae7c66e609fdcf12cf"}, +] + +[package.dependencies] +anyio = {version = ">=4.10", markers = "sys_platform != \"emscripten\""} +httpcore2 = {version = "2.12.0", markers = "sys_platform != \"emscripten\""} +httpx2-jsfetch = {version = "*", markers = "sys_platform == \"emscripten\" and python_version >= \"3.12\""} +idna = ">=3.18" +truststore = {version = ">=0.10", markers = "sys_platform != \"emscripten\""} +typing-extensions = {version = ">=4.5.0", markers = "python_version < \"3.13\""} + +[package.extras] +brotli = ["brotli (>=1.2.0) ; platform_python_implementation == \"CPython\"", "brotlicffi (>=1.2.0.0) ; platform_python_implementation != \"CPython\""] +cli = ["click (>=8.4.2)", "pygments (==2.*)", "rich (>=10,<16)"] +http2 = ["h2 (>=3,<5)"] +socks = ["socksio (==1.*)"] +ws = ["wsproto (>=1.2)"] +zstd = ["backports-zstd (>=1.0.0) ; python_version <= \"3.13\""] + +[[package]] +name = "httpx2-jsfetch" +version = "1.0" +description = "httpx2 transports for Emscripten/Pyodide, backed by the JavaScript fetch API." +optional = false +python-versions = ">=3.12" +groups = ["main"] +markers = "sys_platform == \"emscripten\" and python_version >= \"3.12\"" +files = [ + {file = "httpx2_jsfetch-1.0-py3-none-any.whl", hash = "sha256:cb916b707601e69a07721aabc8f3f6659be3a6893bc1ff5c6f9e02241df2da32"}, + {file = "httpx2_jsfetch-1.0.tar.gz", hash = "sha256:70a0e3eabfef7cce5ad9c629f7d01ca05e418f586646f4ddf14782e4c1454c60"}, +] + [[package]] name = "huggingface-hub" version = "1.28.0" @@ -1931,7 +1995,6 @@ files = [ {file = "python-dateutil-2.9.0.post0.tar.gz", hash = "sha256:37dd54208da7e1cd875388217d5e00ebd4179249f90fb72437e91a35459a0ad3"}, {file = "python_dateutil-2.9.0.post0-py2.py3-none-any.whl", hash = "sha256:a8b2bc7bffae282281c8140a97d3aa9c14da0b136dfe83f850eea9a5f7470427"}, ] -markers = {main = "extra == \"oci\""} [package.dependencies] six = ">=1.5" @@ -2093,7 +2156,6 @@ files = [ {file = "six-1.17.0-py2.py3-none-any.whl", hash = "sha256:4721f391ed90541fddacab5acf947aa0d3dc7d27b2e1e8eda2be8970586c3274"}, {file = "six-1.17.0.tar.gz", hash = "sha256:ff70335d468e7eb6ec65b95b99d3a2836546063f63acc5171de367e834932a81"}, ] -markers = {main = "extra == \"oci\""} [[package]] name = "tokenizers" @@ -2209,6 +2271,19 @@ notebook = ["ipywidgets (>=6)"] slack = ["envwrap", "slack-sdk"] telegram = ["envwrap", "requests"] +[[package]] +name = "truststore" +version = "0.10.4" +description = "Verify certificates using native system trust stores" +optional = false +python-versions = ">=3.10" +groups = ["main"] +markers = "sys_platform != \"emscripten\"" +files = [ + {file = "truststore-0.10.4-py3-none-any.whl", hash = "sha256:adaeaecf1cbb5f4de3b1959b42d41f6fab57b2b1666adb59e89cb0b53361d981"}, + {file = "truststore-0.10.4.tar.gz", hash = "sha256:9d91bd436463ad5e4ee4aba766628dd6cd7010cf3e2461756b3303710eebc301"}, +] + [[package]] name = "types-python-dateutil" version = "2.9.0.20260807" @@ -2402,10 +2477,10 @@ multidict = ">=4.0" propcache = ">=0.2.1" [extras] -aiohttp = ["aiohttp", "httpx-aiohttp"] +aiohttp = ["aiohttp", "httpx", "httpx-aiohttp"] oci = ["oci"] [metadata] lock-version = "2.1" python-versions = "^3.10" -content-hash = "ed02357a175d075bec4b85a4af74f38b816ebd7b1b11cc8fba2015a53d50a2c9" +content-hash = "afb7a50736d983256c8119a294bd5870e806abe58bf9975b00ec0cc52dcfad8a" diff --git a/pyproject.toml b/pyproject.toml index dc71b35c7..5822a7a78 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -40,7 +40,8 @@ Repository = 'https://github.com/cohere-ai/cohere-python' python = "^3.10" aiohttp = { version = ">=3.14.1,<4", optional = true, python = ">=3.10"} fastavro = "^1.9.4" -httpx = ">=0.25.0" +httpx = { version = ">=0.25.0", optional = true } # only used by the aiohttp extra (httpx-aiohttp) +httpx2 = "^2.12.0" httpx-aiohttp = { version = "^0.1.8", optional = true, python = ">=3.10"} oci = { version = "^2.165.0", optional = true} pydantic = ">= 1.9.2" @@ -101,4 +102,4 @@ build-backend = "poetry.core.masonry.api" [tool.poetry.extras] oci=["oci"] -aiohttp=["aiohttp", "httpx-aiohttp"] +aiohttp=["aiohttp", "httpx-aiohttp", "httpx"] diff --git a/src/cohere/_default_clients.py b/src/cohere/_default_clients.py index 4c23bb335..a55830dcd 100644 --- a/src/cohere/_default_clients.py +++ b/src/cohere/_default_clients.py @@ -2,7 +2,7 @@ import typing -import httpx +import httpx2 SDK_DEFAULT_TIMEOUT = 60 @@ -10,7 +10,7 @@ import httpx_aiohttp # type: ignore[import-not-found] except ImportError: - class DefaultAioHttpClient(httpx.AsyncClient): # type: ignore + class DefaultAioHttpClient(httpx2.AsyncClient): # type: ignore def __init__(self, **kwargs: typing.Any) -> None: raise RuntimeError("To use the aiohttp client, install the aiohttp extra: pip install cohere[aiohttp]") @@ -23,7 +23,7 @@ def __init__(self, **kwargs: typing.Any) -> None: super().__init__(**kwargs) -class DefaultAsyncHttpxClient(httpx.AsyncClient): +class DefaultAsyncHttpxClient(httpx2.AsyncClient): def __init__(self, **kwargs: typing.Any) -> None: kwargs.setdefault("timeout", SDK_DEFAULT_TIMEOUT) kwargs.setdefault("follow_redirects", True) diff --git a/src/cohere/aws_client.py b/src/cohere/aws_client.py index 12a168276..606b19bc6 100644 --- a/src/cohere/aws_client.py +++ b/src/cohere/aws_client.py @@ -3,8 +3,8 @@ import re import typing -import httpx -from httpx import URL, SyncByteStream, ByteStream +import httpx2 +from httpx2 import URL, SyncByteStream, ByteStream from . import GenerateStreamedResponse, Generation, \ NonStreamedChatResponse, EmbedResponse, StreamedChatResponse, RerankResponse, ApiMeta, ApiMetaTokens, \ @@ -32,7 +32,7 @@ def __init__( client_name="n/a", timeout=timeout, api_key="n/a", - httpx_client=httpx.Client( + httpx_client=httpx2.Client( event_hooks=get_event_hooks( service=service, aws_access_key=aws_access_key, @@ -63,7 +63,7 @@ def __init__( client_name="n/a", timeout=timeout, api_key="n/a", - httpx_client=httpx.Client( + httpx_client=httpx2.Client( event_hooks=get_event_hooks( service=service, aws_access_key=aws_access_key, @@ -135,7 +135,7 @@ def __iter__(self) -> typing.Iterator[bytes]: } -def stream_generator(response: httpx.Response, endpoint: str) -> typing.Iterator[bytes]: +def stream_generator(response: httpx2.Response, endpoint: str) -> typing.Iterator[bytes]: regex = r"{[^\}]*}" for _text in response.iter_lines(): @@ -152,7 +152,7 @@ def stream_generator(response: httpx.Response, endpoint: str) -> typing.Iterator yield (json.dumps(parsed.dict()) + "\n").encode("utf-8") # type: ignore -def map_token_counts(response: httpx.Response) -> ApiMeta: +def map_token_counts(response: httpx2.Response) -> ApiMeta: input_tokens = int(response.headers.get("X-Amzn-Bedrock-Input-Token-Count", -1)) output_tokens = int(response.headers.get("X-Amzn-Bedrock-Output-Token-Count", -1)) return ApiMeta( @@ -163,14 +163,14 @@ def map_token_counts(response: httpx.Response) -> ApiMeta: def map_response_from_bedrock(): def _hook( - response: httpx.Response, + response: httpx2.Response, ) -> None: stream = response.headers["content-type"] == "application/vnd.amazon.eventstream" endpoint = response.request.extensions["endpoint"] output: typing.Iterator[bytes] if stream: - output = stream_generator(httpx.Response( + output = stream_generator(httpx2.Response( stream=response.stream, status_code=response.status_code, ), endpoint) @@ -221,7 +221,7 @@ def map_request_to_bedrock( credentials = session.get_credentials() signer = lazy_botocore().auth.SigV4Auth(credentials, service, aws_region) - def _event_hook(request: httpx.Request) -> None: + def _event_hook(request: httpx2.Request) -> None: headers = request.headers.copy() del headers["connection"] @@ -263,7 +263,7 @@ def _event_hook(request: httpx.Request) -> None: ) signer.add_auth(aws_request) - request.headers = httpx.Headers(aws_request.prepare().headers) + request.headers = httpx2.Headers(aws_request.prepare().headers) request.extensions["endpoint"] = endpoint return _event_hook diff --git a/src/cohere/base_client.py b/src/cohere/base_client.py index d34ef0581..7da3e5b41 100644 --- a/src/cohere/base_client.py +++ b/src/cohere/base_client.py @@ -5,7 +5,7 @@ import os import typing -import httpx +import httpx2 from .core.api_error import ApiError from .core.client_wrapper import AsyncClientWrapper, SyncClientWrapper from .core.http_client import get_keepalive_socket_options @@ -87,7 +87,7 @@ class BaseCohere: Additional headers to send with every request. timeout : typing.Optional[float] - The timeout to be used, in seconds, for requests. By default the timeout is 300 seconds, unless a custom httpx client is used, in which case this default is not enforced. + The timeout to be used, in seconds, for requests. By default the timeout is 300 seconds, unless a custom httpx2 client is used, in which case this default is not enforced. max_retries : typing.Optional[int] The default maximum number of retries for failed requests. Defaults to 2. Per-request `max_retries` in `request_options` takes precedence over this value. @@ -99,10 +99,10 @@ class BaseCohere: The maximum number of reconnection attempts for resumable streaming endpoints. Defaults to no limit. Per-request `max_stream_reconnection_attempts` in `request_options` takes precedence over this value. follow_redirects : typing.Optional[bool] - Whether the default httpx client follows redirects or not, this is irrelevant if a custom httpx client is passed in. + Whether the default httpx2 client follows redirects or not, this is irrelevant if a custom httpx2 client is passed in. - httpx_client : typing.Optional[httpx.Client] - The httpx client to use for making requests, a preconfigured client is used by default, however this is useful should you want to pass in any custom httpx configuration. + httpx_client : typing.Optional[httpx2.Client] + The httpx2 client to use for making requests, a preconfigured client is used by default, however this is useful should you want to pass in any custom httpx2 configuration. logging : typing.Optional[typing.Union[LogConfig, Logger]] Configure logging for the SDK. Accepts a LogConfig dict with 'level' (debug/info/warn/error), 'logger' (custom logger implementation), and 'silent' (boolean, defaults to True) fields. You can also pass a pre-configured Logger instance. @@ -130,7 +130,7 @@ def __init__( stream_reconnection_enabled: typing.Optional[bool] = None, max_stream_reconnection_attempts: typing.Optional[int] = None, follow_redirects: typing.Optional[bool] = True, - httpx_client: typing.Optional[httpx.Client] = None, + httpx_client: typing.Optional[httpx2.Client] = None, logging: typing.Optional[typing.Union[LogConfig, Logger]] = None, ): _defaulted_timeout = timeout if timeout is not None else 300 if httpx_client is None else None @@ -144,15 +144,15 @@ def __init__( headers=headers, httpx_client=httpx_client if httpx_client is not None - else httpx.Client( + else httpx2.Client( timeout=_defaulted_timeout, follow_redirects=follow_redirects, - transport=httpx.HTTPTransport(socket_options=get_keepalive_socket_options(idle=60, intvl=30, cnt=5)), + transport=httpx2.HTTPTransport(socket_options=get_keepalive_socket_options(idle=60, intvl=30, cnt=5)), ) if follow_redirects is not None - else httpx.Client( + else httpx2.Client( timeout=_defaulted_timeout, - transport=httpx.HTTPTransport(socket_options=get_keepalive_socket_options(idle=60, intvl=30, cnt=5)), + transport=httpx2.HTTPTransport(socket_options=get_keepalive_socket_options(idle=60, intvl=30, cnt=5)), ), timeout=_defaulted_timeout, max_retries=_defaulted_max_retries, @@ -1610,10 +1610,10 @@ def audio(self): def _make_default_async_client( timeout: typing.Optional[float], follow_redirects: typing.Optional[bool], - transport: typing.Optional[httpx.AsyncBaseTransport] = None, -) -> httpx.AsyncClient: + transport: typing.Optional[httpx2.AsyncBaseTransport] = None, +) -> httpx2.AsyncClient: if transport is None: - transport = httpx.AsyncHTTPTransport(socket_options=get_keepalive_socket_options(idle=60, intvl=30, cnt=5)) + transport = httpx2.AsyncHTTPTransport(socket_options=get_keepalive_socket_options(idle=60, intvl=30, cnt=5)) try: import httpx_aiohttp # type: ignore[import-not-found] except ImportError: @@ -1624,8 +1624,8 @@ def _make_default_async_client( return httpx_aiohttp.HttpxAiohttpClient(timeout=timeout) if follow_redirects is not None: - return httpx.AsyncClient(timeout=timeout, follow_redirects=follow_redirects, transport=transport) - return httpx.AsyncClient(timeout=timeout, transport=transport) + return httpx2.AsyncClient(timeout=timeout, follow_redirects=follow_redirects, transport=transport) + return httpx2.AsyncClient(timeout=timeout, transport=transport) class AsyncBaseCohere: @@ -1655,7 +1655,7 @@ class AsyncBaseCohere: An async callable that returns a bearer token. Use this when token acquisition involves async I/O (e.g., refreshing tokens via an async HTTP client). When provided, this is used instead of the synchronous token for async requests. timeout : typing.Optional[float] - The timeout to be used, in seconds, for requests. By default the timeout is 300 seconds, unless a custom httpx client is used, in which case this default is not enforced. + The timeout to be used, in seconds, for requests. By default the timeout is 300 seconds, unless a custom httpx2 client is used, in which case this default is not enforced. max_retries : typing.Optional[int] The default maximum number of retries for failed requests. Defaults to 2. Per-request `max_retries` in `request_options` takes precedence over this value. @@ -1667,10 +1667,10 @@ class AsyncBaseCohere: The maximum number of reconnection attempts for resumable streaming endpoints. Defaults to no limit. Per-request `max_stream_reconnection_attempts` in `request_options` takes precedence over this value. follow_redirects : typing.Optional[bool] - Whether the default httpx client follows redirects or not, this is irrelevant if a custom httpx client is passed in. + Whether the default httpx2 client follows redirects or not, this is irrelevant if a custom httpx2 client is passed in. - httpx_client : typing.Optional[httpx.AsyncClient] - The httpx client to use for making requests, a preconfigured client is used by default, however this is useful should you want to pass in any custom httpx configuration. + httpx_client : typing.Optional[httpx2.AsyncClient] + The httpx2 client to use for making requests, a preconfigured client is used by default, however this is useful should you want to pass in any custom httpx2 configuration. logging : typing.Optional[typing.Union[LogConfig, Logger]] Configure logging for the SDK. Accepts a LogConfig dict with 'level' (debug/info/warn/error), 'logger' (custom logger implementation), and 'silent' (boolean, defaults to True) fields. You can also pass a pre-configured Logger instance. @@ -1699,7 +1699,7 @@ def __init__( stream_reconnection_enabled: typing.Optional[bool] = None, max_stream_reconnection_attempts: typing.Optional[int] = None, follow_redirects: typing.Optional[bool] = True, - httpx_client: typing.Optional[httpx.AsyncClient] = None, + httpx_client: typing.Optional[httpx2.AsyncClient] = None, logging: typing.Optional[typing.Union[LogConfig, Logger]] = None, ): _defaulted_timeout = timeout if timeout is not None else 300 if httpx_client is None else None diff --git a/src/cohere/client.py b/src/cohere/client.py index dfb7e59f7..2630f4299 100644 --- a/src/cohere/client.py +++ b/src/cohere/client.py @@ -5,7 +5,7 @@ from tokenizers import Tokenizer # type: ignore import logging -import httpx +import httpx2 from cohere.types.detokenize_response import DetokenizeResponse from cohere.types.tokenize_response import TokenizeResponse @@ -140,7 +140,7 @@ def __init__( client_name: typing.Optional[str] = None, timeout: typing.Optional[float] = None, max_retries: typing.Optional[int] = None, - httpx_client: typing.Optional[httpx.Client] = None, + httpx_client: typing.Optional[httpx2.Client] = None, thread_pool_executor: ThreadPoolExecutor = ThreadPoolExecutor(64), log_warning_experimental_features: bool = True, ): @@ -389,7 +389,7 @@ def __init__( client_name: typing.Optional[str] = None, timeout: typing.Optional[float] = None, max_retries: typing.Optional[int] = None, - httpx_client: typing.Optional[httpx.AsyncClient] = None, + httpx_client: typing.Optional[httpx2.AsyncClient] = None, thread_pool_executor: ThreadPoolExecutor = ThreadPoolExecutor(64), log_warning_experimental_features: bool = True, ): diff --git a/src/cohere/client_v2.py b/src/cohere/client_v2.py index 5a3633bbd..76eb9e1e2 100644 --- a/src/cohere/client_v2.py +++ b/src/cohere/client_v2.py @@ -2,7 +2,7 @@ import typing from concurrent.futures import ThreadPoolExecutor -import httpx +import httpx2 from .client import AsyncClient, Client from .environment import ClientEnvironment from .v2.client import AsyncRawV2Client, AsyncV2Client, RawV2Client, V2Client @@ -39,7 +39,7 @@ def __init__( client_name: typing.Optional[str] = None, timeout: typing.Optional[float] = None, max_retries: typing.Optional[int] = None, - httpx_client: typing.Optional[httpx.Client] = None, + httpx_client: typing.Optional[httpx2.Client] = None, thread_pool_executor: ThreadPoolExecutor = ThreadPoolExecutor(64), log_warning_experimental_features: bool = True, ): @@ -74,7 +74,7 @@ def __init__( client_name: typing.Optional[str] = None, timeout: typing.Optional[float] = None, max_retries: typing.Optional[int] = None, - httpx_client: typing.Optional[httpx.AsyncClient] = None, + httpx_client: typing.Optional[httpx2.AsyncClient] = None, thread_pool_executor: ThreadPoolExecutor = ThreadPoolExecutor(64), log_warning_experimental_features: bool = True, ): diff --git a/src/cohere/core/client_wrapper.py b/src/cohere/core/client_wrapper.py index f2ed6c3fb..17f1caf1f 100644 --- a/src/cohere/core/client_wrapper.py +++ b/src/cohere/core/client_wrapper.py @@ -2,7 +2,7 @@ import typing -import httpx +import httpx2 from .http_client import AsyncHttpClient, HttpClient from .logging import LogConfig, Logger @@ -86,7 +86,7 @@ def __init__( stream_reconnection_enabled: typing.Optional[bool] = None, max_stream_reconnection_attempts: typing.Optional[int] = None, logging: typing.Optional[typing.Union[LogConfig, Logger]] = None, - httpx_client: httpx.Client, + httpx_client: httpx2.Client, ): super().__init__( client_name=client_name, @@ -123,7 +123,7 @@ def __init__( max_stream_reconnection_attempts: typing.Optional[int] = None, logging: typing.Optional[typing.Union[LogConfig, Logger]] = None, async_token: typing.Optional[typing.Callable[[], typing.Awaitable[str]]] = None, - httpx_client: httpx.AsyncClient, + httpx_client: httpx2.AsyncClient, ): super().__init__( client_name=client_name, diff --git a/src/cohere/core/file.py b/src/cohere/core/file.py index 44b0d27c0..90fcf9c69 100644 --- a/src/cohere/core/file.py +++ b/src/cohere/core/file.py @@ -2,8 +2,8 @@ from typing import IO, Dict, List, Mapping, Optional, Tuple, Union, cast -# File typing inspired by the flexibility of types within the httpx library -# https://github.com/encode/httpx/blob/master/httpx/_types.py +# File typing inspired by the flexibility of types within the httpx2 library +# https://github.com/encode/httpx2/blob/master/httpx2/_types.py FileContent = Union[IO[bytes], bytes, str] File = Union[ # file (or bytes) @@ -30,7 +30,7 @@ def convert_file_dict_to_httpx_tuples( name of the file and the second is the file object. Typically HTTPX wants a dict, but to be able to send lists of files, you have to use the list approach (which also works for non-lists) - https://github.com/encode/httpx/pull/1032 + https://github.com/encode/httpx2/pull/1032 """ httpx_tuples = [] diff --git a/src/cohere/core/http_client.py b/src/cohere/core/http_client.py index 409d26709..0949132e9 100644 --- a/src/cohere/core/http_client.py +++ b/src/cohere/core/http_client.py @@ -9,7 +9,17 @@ from contextlib import asynccontextmanager, contextmanager from random import random -import httpx +import httpx2 + +# The optional aiohttp extra builds its async client on httpx 0.x (httpx-aiohttp), +# so that path raises httpx exceptions rather than httpx2 ones. Catch both. +try: + import httpx as _httpx_aiohttp + + _AIOHTTP_CONNECT_ERRORS: tuple = (_httpx_aiohttp.ConnectError, _httpx_aiohttp.RemoteProtocolError) +except ImportError: + _AIOHTTP_CONNECT_ERRORS = () + from .file import File, convert_file_dict_to_httpx_tuples from .force_multipart import FORCE_MULTIPART from .jsonable_encoder import jsonable_encoder @@ -17,7 +27,7 @@ from .query_encoder import encode_query from .remove_none_from_dict import remove_none_from_dict as remove_none_from_dict from .request_options import RequestOptions -from httpx._types import RequestFiles +from httpx2._types import RequestFiles INITIAL_RETRY_DELAY_SECONDS = 1.0 MAX_RETRY_DELAY_SECONDS = 60.0 @@ -42,8 +52,8 @@ def get_keepalive_socket_options( Windows, but ``TCP_KEEPALIVE`` on macOS. - ``TCP_KEEPINTVL`` / ``TCP_KEEPCNT`` exist on Linux/macOS/modern Windows. - Passing these tuples to ``httpx.HTTPTransport(socket_options=...)`` / - ``httpx.AsyncHTTPTransport(socket_options=...)`` applies them to every + Passing these tuples to ``httpx2.HTTPTransport(socket_options=...)`` / + ``httpx2.AsyncHTTPTransport(socket_options=...)`` applies them to every connection the transport opens. """ opts: typing.List[typing.Tuple[int, int, int]] = [(socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1)] @@ -57,7 +67,7 @@ def get_keepalive_socket_options( return opts -def _parse_retry_after(response_headers: httpx.Headers) -> typing.Optional[float]: +def _parse_retry_after(response_headers: httpx2.Headers) -> typing.Optional[float]: """ This function parses the `Retry-After` header in a HTTP response and returns the number of seconds to wait. @@ -110,7 +120,7 @@ def _add_symmetric_jitter(delay: float) -> float: return delay * jitter_multiplier -def _parse_x_ratelimit_reset(response_headers: httpx.Headers) -> typing.Optional[float]: +def _parse_x_ratelimit_reset(response_headers: httpx2.Headers) -> typing.Optional[float]: """ Parse the X-RateLimit-Reset header (Unix timestamp in seconds). Returns seconds to wait, or None if header is missing/invalid. @@ -130,7 +140,7 @@ def _parse_x_ratelimit_reset(response_headers: httpx.Headers) -> typing.Optional return None -def _retry_timeout(response: httpx.Response, retries: int) -> float: +def _retry_timeout(response: httpx2.Response, retries: int) -> float: """ Determine the amount of time to wait before retrying a request. This function begins by trying to parse a retry-after header from the response, and then proceeds to use exponential backoff @@ -158,7 +168,7 @@ def _retry_timeout_from_retries(retries: int) -> float: return _add_symmetric_jitter(backoff) -def _should_retry(response: httpx.Response) -> bool: +def _should_retry(response: httpx2.Response) -> bool: return response.status_code >= 500 or response.status_code in [429, 408, 409] @@ -219,7 +229,7 @@ def _maybe_filter_none_from_multipart_data( ) -> typing.Optional[typing.Any]: """ Filter None values from data body for multipart/form requests. - This prevents httpx from converting None to empty strings in multipart encoding. + This prevents httpx2 from converting None to empty strings in multipart encoding. Only applies when files are present or force_multipart is True. """ if data is not None and isinstance(data, typing.Mapping) and (request_files or force_multipart): @@ -300,7 +310,7 @@ class HttpClient: def __init__( self, *, - httpx_client: httpx.Client, + httpx_client: httpx2.Client, base_timeout: typing.Callable[[], typing.Optional[float]], base_headers: typing.Callable[[], typing.Dict[str, str]], base_url: typing.Optional[typing.Callable[[], str]] = None, @@ -344,7 +354,7 @@ def request( retries: int = 0, omit: typing.Optional[typing.Any] = None, force_multipart: typing.Optional[bool] = None, - ) -> httpx.Response: + ) -> httpx2.Response: base_url = self.get_base_url(base_url) _timeout = ( request_options.get("timeout") @@ -353,7 +363,7 @@ def request( if request_options is not None and request_options.get("timeout_in_seconds") is not None else self.base_timeout() ) - timeout = _timeout if _timeout is not None else httpx.USE_CLIENT_DEFAULT + timeout = _timeout if _timeout is not None else httpx2.USE_CLIENT_DEFAULT json_body, data_body = get_request_body(json=json, data=data, request_options=request_options, omit=omit) @@ -368,8 +378,8 @@ def request( data_body = _maybe_filter_none_from_multipart_data(data_body, request_files, force_multipart) - # Compute encoded params separately to avoid passing empty list to httpx - # (httpx strips existing query params from URL when params=[] is passed) + # Compute encoded params separately to avoid passing empty list to httpx2 + # (httpx2 strips existing query params from URL when params=[] is passed) _encoded_params = encode_query( jsonable_encoder( remove_none_from_dict( @@ -426,7 +436,7 @@ def request( files=request_files, timeout=timeout, ) - except (httpx.ConnectError, httpx.RemoteProtocolError): + except (httpx2.ConnectError, httpx2.RemoteProtocolError): if retries < max_retries: time.sleep(_retry_timeout_from_retries(retries=retries)) return self.request( @@ -507,7 +517,7 @@ def stream( retries: int = 0, omit: typing.Optional[typing.Any] = None, force_multipart: typing.Optional[bool] = None, - ) -> typing.Iterator[httpx.Response]: + ) -> typing.Iterator[httpx2.Response]: base_url = self.get_base_url(base_url) _timeout = ( request_options.get("timeout") @@ -516,7 +526,7 @@ def stream( if request_options is not None and request_options.get("timeout_in_seconds") is not None else self.base_timeout() ) - timeout = _timeout if _timeout is not None else httpx.USE_CLIENT_DEFAULT + timeout = _timeout if _timeout is not None else httpx2.USE_CLIENT_DEFAULT request_files: typing.Optional[RequestFiles] = ( convert_file_dict_to_httpx_tuples(remove_omit_from_dict(remove_none_from_dict(files), omit)) @@ -531,8 +541,8 @@ def stream( data_body = _maybe_filter_none_from_multipart_data(data_body, request_files, force_multipart) - # Compute encoded params separately to avoid passing empty list to httpx - # (httpx strips existing query params from URL when params=[] is passed) + # Compute encoded params separately to avoid passing empty list to httpx2 + # (httpx2 strips existing query params from URL when params=[] is passed) _encoded_params = encode_query( jsonable_encoder( remove_none_from_dict( @@ -588,7 +598,7 @@ class AsyncHttpClient: def __init__( self, *, - httpx_client: httpx.AsyncClient, + httpx_client: httpx2.AsyncClient, base_timeout: typing.Callable[[], typing.Optional[float]], base_headers: typing.Callable[[], typing.Dict[str, str]], base_url: typing.Optional[typing.Callable[[], str]] = None, @@ -639,7 +649,7 @@ async def request( retries: int = 0, omit: typing.Optional[typing.Any] = None, force_multipart: typing.Optional[bool] = None, - ) -> httpx.Response: + ) -> httpx2.Response: base_url = self.get_base_url(base_url) _timeout = ( request_options.get("timeout") @@ -648,7 +658,7 @@ async def request( if request_options is not None and request_options.get("timeout_in_seconds") is not None else self.base_timeout() ) - timeout = _timeout if _timeout is not None else httpx.USE_CLIENT_DEFAULT + timeout = _timeout if _timeout is not None else httpx2.USE_CLIENT_DEFAULT request_files: typing.Optional[RequestFiles] = ( convert_file_dict_to_httpx_tuples(remove_omit_from_dict(remove_none_from_dict(files), omit)) @@ -666,8 +676,8 @@ async def request( # Get headers (supports async token providers) _headers = await self._get_headers() - # Compute encoded params separately to avoid passing empty list to httpx - # (httpx strips existing query params from URL when params=[] is passed) + # Compute encoded params separately to avoid passing empty list to httpx2 + # (httpx2 strips existing query params from URL when params=[] is passed) _encoded_params = encode_query( jsonable_encoder( remove_none_from_dict( @@ -724,7 +734,7 @@ async def request( files=request_files, timeout=timeout, ) - except (httpx.ConnectError, httpx.RemoteProtocolError): + except (httpx2.ConnectError, httpx2.RemoteProtocolError, *_AIOHTTP_CONNECT_ERRORS): if retries < max_retries: await asyncio.sleep(_retry_timeout_from_retries(retries=retries)) return await self.request( @@ -805,7 +815,7 @@ async def stream( retries: int = 0, omit: typing.Optional[typing.Any] = None, force_multipart: typing.Optional[bool] = None, - ) -> typing.AsyncIterator[httpx.Response]: + ) -> typing.AsyncIterator[httpx2.Response]: base_url = self.get_base_url(base_url) _timeout = ( request_options.get("timeout") @@ -814,7 +824,7 @@ async def stream( if request_options is not None and request_options.get("timeout_in_seconds") is not None else self.base_timeout() ) - timeout = _timeout if _timeout is not None else httpx.USE_CLIENT_DEFAULT + timeout = _timeout if _timeout is not None else httpx2.USE_CLIENT_DEFAULT request_files: typing.Optional[RequestFiles] = ( convert_file_dict_to_httpx_tuples(remove_omit_from_dict(remove_none_from_dict(files), omit)) @@ -832,8 +842,8 @@ async def stream( # Get headers (supports async token providers) _headers = await self._get_headers() - # Compute encoded params separately to avoid passing empty list to httpx - # (httpx strips existing query params from URL when params=[] is passed) + # Compute encoded params separately to avoid passing empty list to httpx2 + # (httpx2 strips existing query params from URL when params=[] is passed) _encoded_params = encode_query( jsonable_encoder( remove_none_from_dict( diff --git a/src/cohere/core/http_response.py b/src/cohere/core/http_response.py index 00bb1096d..c0b5bfa99 100644 --- a/src/cohere/core/http_response.py +++ b/src/cohere/core/http_response.py @@ -2,7 +2,7 @@ from typing import Dict, Generic, TypeVar -import httpx +import httpx2 # Generic to represent the underlying type of the data wrapped by the HTTP response. T = TypeVar("T") @@ -11,9 +11,9 @@ class BaseHttpResponse: """Minimalist HTTP response wrapper that exposes response headers and status code.""" - _response: httpx.Response + _response: httpx2.Response - def __init__(self, response: httpx.Response): + def __init__(self, response: httpx2.Response): self._response = response @property @@ -30,7 +30,7 @@ class HttpResponse(Generic[T], BaseHttpResponse): _data: T - def __init__(self, response: httpx.Response, data: T): + def __init__(self, response: httpx2.Response, data: T): super().__init__(response) self._data = data @@ -47,7 +47,7 @@ class AsyncHttpResponse(Generic[T], BaseHttpResponse): _data: T - def __init__(self, response: httpx.Response, data: T): + def __init__(self, response: httpx2.Response, data: T): super().__init__(response) self._data = data diff --git a/src/cohere/core/http_sse/_api.py b/src/cohere/core/http_sse/_api.py index 9ca5602c2..6706197f5 100644 --- a/src/cohere/core/http_sse/_api.py +++ b/src/cohere/core/http_sse/_api.py @@ -16,7 +16,16 @@ ) import anyio -import httpx +import httpx2 + +# The optional aiohttp extra's async client raises httpx (0.x) transport errors. +try: + import httpx as _httpx_aiohttp + + _AIOHTTP_TRANSPORT_ERRORS: tuple = (_httpx_aiohttp.TransportError,) +except ImportError: + _AIOHTTP_TRANSPORT_ERRORS = () + from ._decoders import SSEDecoder from ._exceptions import SSEError from ._models import ServerSentEvent @@ -31,12 +40,12 @@ # A reconnect callback re-issues the original request (with a ``Last-Event-ID`` # header set to the supplied event id) and returns a *context manager* yielding -# a fresh streaming ``httpx.Response``. Sync clients supply a sync context +# a fresh streaming ``httpx2.Response``. Sync clients supply a sync context # manager; async clients supply an async one. class EventSource: def __init__( self, - response: httpx.Response, + response: httpx2.Response, *, resumable: bool = False, stream_reconnection_enabled: bool = True, @@ -52,7 +61,7 @@ def __init__( self._reconnect = reconnect @staticmethod - def _is_event_stream(response: httpx.Response) -> bool: + def _is_event_stream(response: httpx2.Response) -> bool: content_type = response.headers.get("content-type", "").partition(";")[0] return "text/event-stream" in content_type @@ -63,17 +72,17 @@ def _check_content_type(self) -> None: f"Expected response header Content-Type to contain 'text/event-stream', got {content_type!r}" ) - def _is_reconnect_response_usable(self, response: httpx.Response) -> bool: + def _is_reconnect_response_usable(self, response: httpx2.Response) -> bool: """Whether a reconnected response can be resumed as an SSE stream. - ``httpx.stream`` does not raise on non-success status, so a resume that + ``httpx2.stream`` does not raise on non-success status, so a resume that returns an error page (e.g. ``200 text/html`` or a ``500`` body) would otherwise be parsed as SSE and yield garbage/zero events. Such a response is treated as a failed attempt (back off and retry) instead. """ return response.status_code < 400 and self._is_event_stream(response) - def _get_charset(self, response: Optional[httpx.Response] = None) -> str: + def _get_charset(self, response: Optional[httpx2.Response] = None) -> str: """Extract charset from Content-Type header, fallback to UTF-8.""" resolved = response if response is not None else self._response content_type = resolved.headers.get("content-type", "") @@ -95,7 +104,7 @@ def _get_charset(self, response: Optional[httpx.Response] = None) -> str: return "utf-8" @property - def response(self) -> httpx.Response: + def response(self) -> httpx2.Response: return self._response @staticmethod @@ -110,7 +119,7 @@ def _normalize_sse_line_endings(buf: str) -> str: return buf[:-1].replace("\r", "\n") + "\r" return buf.replace("\r", "\n") - def _new_text_decoder(self, response: Optional[httpx.Response] = None) -> "codecs.IncrementalDecoder": + def _new_text_decoder(self, response: Optional[httpx2.Response] = None) -> "codecs.IncrementalDecoder": return codecs.getincrementaldecoder(self._get_charset(response))(errors="replace") def _reconnect_applicable(self) -> bool: @@ -181,7 +190,7 @@ async def _asleep_before_reconnect(self, last_retry: Optional[int]) -> None: def _decode_response( self, - response: httpx.Response, + response: httpx2.Response, decoder: SSEDecoder, text_decoder: "codecs.IncrementalDecoder", ) -> Iterator[ServerSentEvent]: @@ -205,7 +214,7 @@ def _decode_response( async def _adecode_response( self, - response: httpx.Response, + response: httpx2.Response, decoder: SSEDecoder, text_decoder: "codecs.IncrementalDecoder", ) -> AsyncGenerator[ServerSentEvent, None]: @@ -270,10 +279,10 @@ def iter_sse(self) -> Iterator[ServerSentEvent]: # ``None`` means there is no live stream to read this iteration (e.g. a # failed reconnect); the loop then re-evaluates the reconnect decision # without re-reading an exhausted response. - response: Optional[httpx.Response] = self._response + response: Optional[httpx2.Response] = self._response # Context manager for a response we opened ourselves and must close. # The initial response is owned by the caller, so it starts as None. - owned_cm: Optional[ContextManager[httpx.Response]] = None + owned_cm: Optional[ContextManager[httpx2.Response]] = None try: while True: if response is not None: @@ -287,9 +296,9 @@ def iter_sse(self) -> Iterator[ServerSentEvent]: # A protocol violation (e.g. an oversized line) is a # genuine error, not a dropped connection; propagate it. # Listed first because ``SSEError`` subclasses - # ``httpx.TransportError``. + # ``httpx2.TransportError``. raise - except httpx.TransportError: + except httpx2.TransportError: # A transport error mid-stream (e.g. the server dropped # the connection: ``ReadError``/``RemoteProtocolError``) # is a premature end. Only swallow it when reconnection @@ -326,7 +335,7 @@ def iter_sse(self) -> Iterator[ServerSentEvent]: assert self._reconnect is not None # guaranteed by _should_reconnect try: - cm: ContextManager[httpx.Response] = self._reconnect(last_dispatched_id or "") + cm: ContextManager[httpx2.Response] = self._reconnect(last_dispatched_id or "") new_response = cm.__enter__() except Exception: # A failed reconnect consumes an attempt; back off and retry. @@ -358,8 +367,8 @@ async def aiter_sse(self) -> AsyncGenerator[ServerSentEvent, None]: last_retry: Optional[int] = None reconnect_attempts = 0 - response: Optional[httpx.Response] = self._response - owned_cm: Optional[AsyncContextManager[httpx.Response]] = None + response: Optional[httpx2.Response] = self._response + owned_cm: Optional[AsyncContextManager[httpx2.Response]] = None try: while True: if response is not None: @@ -373,9 +382,9 @@ async def aiter_sse(self) -> AsyncGenerator[ServerSentEvent, None]: # A protocol violation (e.g. an oversized line) is a # genuine error, not a dropped connection; propagate it. # Listed first because ``SSEError`` subclasses - # ``httpx.TransportError``. + # ``httpx2.TransportError``. raise - except httpx.TransportError: + except (httpx2.TransportError, *_AIOHTTP_TRANSPORT_ERRORS): # A transport error mid-stream (e.g. the server dropped # the connection: ``ReadError``/``RemoteProtocolError``) # is a premature end. Only swallow it when reconnection @@ -408,7 +417,7 @@ async def aiter_sse(self) -> AsyncGenerator[ServerSentEvent, None]: assert self._reconnect is not None # guaranteed by _should_reconnect try: - cm: AsyncContextManager[httpx.Response] = self._reconnect(last_dispatched_id or "") + cm: AsyncContextManager[httpx2.Response] = self._reconnect(last_dispatched_id or "") new_response = await cm.__aenter__() except Exception: response = None @@ -431,7 +440,7 @@ async def aiter_sse(self) -> AsyncGenerator[ServerSentEvent, None]: @contextmanager -def connect_sse(client: httpx.Client, method: str, url: str, **kwargs: Any) -> Iterator[EventSource]: +def connect_sse(client: httpx2.Client, method: str, url: str, **kwargs: Any) -> Iterator[EventSource]: headers = kwargs.pop("headers", {}) headers["Accept"] = "text/event-stream" headers["Cache-Control"] = "no-store" @@ -442,7 +451,7 @@ def connect_sse(client: httpx.Client, method: str, url: str, **kwargs: Any) -> I @asynccontextmanager async def aconnect_sse( - client: httpx.AsyncClient, + client: httpx2.AsyncClient, method: str, url: str, **kwargs: Any, diff --git a/src/cohere/core/http_sse/_exceptions.py b/src/cohere/core/http_sse/_exceptions.py index 81605a8a6..f357eb99f 100644 --- a/src/cohere/core/http_sse/_exceptions.py +++ b/src/cohere/core/http_sse/_exceptions.py @@ -1,7 +1,7 @@ # This file was auto-generated by Fern from our API Definition. -import httpx +import httpx2 -class SSEError(httpx.TransportError): +class SSEError(httpx2.TransportError): pass diff --git a/src/cohere/core/request_options.py b/src/cohere/core/request_options.py index caa6f669b..658332fec 100644 --- a/src/cohere/core/request_options.py +++ b/src/cohere/core/request_options.py @@ -26,7 +26,7 @@ class RequestOptions(typing.TypedDict, total=False): - additional_body_parameters: typing.Dict[str, typing.Any]. A dictionary containing additional parameters to spread into the request's body parameters dict - - chunk_size: int. The size, in bytes, to process each chunk of data being streamed back within the response. This equates to leveraging `chunk_size` within `requests` or `httpx`, and is only leveraged for file downloads. + - chunk_size: int. The size, in bytes, to process each chunk of data being streamed back within the response. This equates to leveraging `chunk_size` within `requests` or `httpx2`, and is only leveraged for file downloads. """ timeout: NotRequired[float] diff --git a/src/cohere/oci_client.py b/src/cohere/oci_client.py index b71415fce..8276079eb 100644 --- a/src/cohere/oci_client.py +++ b/src/cohere/oci_client.py @@ -7,13 +7,13 @@ import typing import uuid -import httpx +import httpx2 import requests from .client import Client, ClientEnvironment from .client_v2 import ClientV2 from .aws_client import Streamer from .manually_maintained.lazy_oci_deps import lazy_oci -from httpx import URL, ByteStream +from httpx2 import URL, ByteStream class OciClient(Client): @@ -90,7 +90,7 @@ def __init__( client_name="n/a", timeout=timeout, api_key="n/a", - httpx_client=httpx.Client( + httpx_client=httpx2.Client( event_hooks=get_event_hooks( oci_config=oci_config, oci_region=oci_region, @@ -204,7 +204,7 @@ def __init__( if oci_region is None: raise ValueError("oci_region must be provided either directly or in OCI config file") - # Create httpx client with OCI event hooks + # Create httpx2 client with OCI event hooks ClientV2.__init__( self, base_url="https://api.cohere.com", # Unused, OCI URL set in hooks @@ -212,7 +212,7 @@ def __init__( client_name="n/a", timeout=timeout, api_key="n/a", - httpx_client=httpx.Client( + httpx_client=httpx2.Client( event_hooks=get_event_hooks( oci_config=oci_config, oci_region=oci_region, @@ -350,7 +350,7 @@ def get_event_hooks( is_v2_client: bool = False, ) -> typing.Dict[str, typing.List[EventHook]]: """ - Create httpx event hooks for OCI request/response transformation. + Create httpx2 event hooks for OCI request/response transformation. Args: oci_config: OCI configuration dictionary @@ -359,7 +359,7 @@ def get_event_hooks( is_v2_client: Whether this is for OciClientV2 (True) or OciClient (False) Returns: - Dictionary of event hooks for httpx + Dictionary of event hooks for httpx2 """ return { "request": [ @@ -390,7 +390,7 @@ def map_request_to_oci( is_v2_client: Whether this is for OciClientV2 (True) or OciClient (False) Returns: - Event hook function for httpx + Event hook function for httpx2 """ oci = lazy_oci() @@ -455,7 +455,7 @@ def __getattr__(self, name: str) -> typing.Any: "session-based authentication, or provide direct credentials via oci_user_id parameter." ) - def _event_hook(request: httpx.Request) -> None: + def _event_hook(request: httpx2.Request) -> None: # Extract Cohere API details path_parts = request.url.path.split("/") endpoint = path_parts[-1] @@ -496,9 +496,9 @@ def _event_hook(request: httpx.Request) -> None: # Sign the request using OCI signer (modifies headers in place) signer.do_request_sign(prepped_request) - # Update httpx request with signed headers + # Update httpx2 request with signed headers request.url = URL(url) - request.headers = httpx.Headers(prepped_request.headers) + request.headers = httpx2.Headers(prepped_request.headers) request.stream = ByteStream(oci_body_bytes) request._content = oci_body_bytes request.extensions["endpoint"] = endpoint @@ -513,10 +513,10 @@ def map_response_from_oci() -> EventHook: Create event hook that transforms OCI responses to Cohere format. Returns: - Event hook function for httpx + Event hook function for httpx2 """ - def _hook(response: httpx.Response) -> None: + def _hook(response: httpx2.Response) -> None: endpoint = response.request.extensions["endpoint"] is_stream = response.request.extensions.get("is_stream", False) is_v2 = response.request.extensions.get("is_v2", False) diff --git a/tests/test_aws_client_unit.py b/tests/test_aws_client_unit.py index 94e584922..04559e847 100644 --- a/tests/test_aws_client_unit.py +++ b/tests/test_aws_client_unit.py @@ -13,7 +13,10 @@ import unittest from unittest.mock import MagicMock, patch -import httpx +try: + import httpx2 as httpx2 +except ImportError: + import httpx2 from cohere.manually_maintained.cohere_aws.mode import Mode @@ -54,7 +57,7 @@ def capture_aws_request(**kwargs): # type: ignore hook = map_request_to_bedrock(service="bedrock", aws_region="us-east-1") - request = httpx.Request( + request = httpx2.Request( method="POST", url="https://api.cohere.com/v1/chat", headers={"connection": "keep-alive"}, diff --git a/tests/test_bedrock_client.py b/tests/test_bedrock_client.py index ee69d8d8e..66a3ecca4 100644 --- a/tests/test_bedrock_client.py +++ b/tests/test_bedrock_client.py @@ -126,7 +126,7 @@ def test_chat_stream(self) -> None: @unittest.skipIf(None == os.getenv("TEST_AWS"), "tests skipped because TEST_AWS is not set") class TestBedrockClientV2(unittest.TestCase): - """Integration tests for BedrockClientV2 (httpx-based). + """Integration tests for BedrockClientV2 (httpx2-based). Fix 1 validation: If these pass, SigV4 signing uses the correct host header, since the request would fail with a signature mismatch otherwise.