Files
Alberto Schiabel b2f5098b28 fix(py): make HTTP status errors catchable as ComposioError on client 2.0 (#4556)
This PR:

- relates to #4537 and ports
https://github.com/ComposioHQ/composio/pull/4543 (which targets `next`
on `composio-client==1.43.0`) to `main`, which pins
`composio-client==2.0.0rc7`
- the 2.0 client has no `_make_status_error` hook; it raises the
module-level `status_error()` from `_decode()` and
`_process_response()`, so `HttpClient` overrides those two instead
- raised errors subclass both the generated class and
`composio.exceptions.ComposioError` (one cached class per generated
class), so `except ComposioError` and `except
composio_client.AuthenticationError` both keep working, with
`status_code`, `response` and `body` unchanged
- the classes support `pickle`/`deepcopy` via `__reduce__`
- adds `python/tests/test_client_errors.py` (from #4543, minus the
`ToolNotFoundError` case, which only exists on `next`): 17 tests, 15 of
which fail without the override
- out of scope: `APIConnectionError`/`APITimeoutError` are still not
`ComposioError`
- when `next` and `main` reconcile, keep this version and drop the
`_make_status_error` override from #4543

Verified locally: full Python suite 1939 passed / 51 skipped (baseline
1922 / 51, no new failures); ruff clean; mypy shows only the 4 errors
already on `main` (`schema_converter.py`, `url_safety.py`).
2026-09-21 17:57:52 +04:00

424 lines
15 KiB
Python

"""
This module is a light wrapper around the auto-generated composio client.
"""
import contextvars
import logging
import os
import platform
import typing as t
from importlib.metadata import version
from uuid import uuid4
import typing_extensions as te
from composio_client import (
DEFAULT_MAX_RETRIES,
NOT_GIVEN,
APIError,
APIStatusError,
NotGiven,
)
from composio_client import Composio as BaseComposio
from httpx import URL, Client, Request, Response, Timeout
from composio.exceptions import ComposioError
from composio.utils.logging import LogLevel, WithLogger, _VerbosityWrapper
ComposioAPIError = APIError
APIEnvironment = te.Literal["production", "staging", "local"]
_SDK_ERROR_CLASSES: t.Dict[t.Type[APIStatusError], t.Type[APIStatusError]] = {}
def _with_sdk_error_base(
error_class: t.Type[APIStatusError],
) -> t.Type[APIStatusError]:
"""
Return a subclass of a generated-client status error that also derives
from the SDK's ``ComposioError``.
The generated client has its own exception root, unrelated to
``composio.exceptions.ComposioError``, so HTTP failures such as an invalid
API key used to escape ``except ComposioError``. Keeping the generated
class as the first base preserves its constructor, ``status_code`` and
``isinstance`` checks, so existing ``except APIStatusError`` handlers keep
working unchanged.
"""
cached = _SDK_ERROR_CLASSES.get(error_class)
if cached is not None:
return cached
sdk_class = t.cast(
t.Type[APIStatusError],
type(
error_class.__name__,
(error_class, ComposioError),
{
"__module__": error_class.__module__,
"__reduce__": _reduce_sdk_error,
},
),
)
# setdefault keeps the first class if two threads race to build one.
return _SDK_ERROR_CLASSES.setdefault(error_class, sdk_class)
def _reduce_sdk_error(self: APIStatusError) -> t.Tuple[t.Any, ...]:
"""
Pickle support for the classes built by ``_with_sdk_error_base``.
They are not module attributes, so pickle cannot find them by name. Record
the generated class instead and rebuild the SDK subclass on load.
"""
error_class = type(self).__mro__[1]
return (
_rebuild_sdk_error,
(error_class, self.message, self.response, self.body),
self.__dict__,
)
def _rebuild_sdk_error(
error_class: t.Type[APIStatusError],
message: str,
response: Response,
body: object,
) -> APIStatusError:
return _with_sdk_error_base(error_class)(message, response=response, body=body)
def _as_sdk_error(error: APIStatusError) -> APIStatusError:
if isinstance(error, ComposioError):
return error
sdk_error = _with_sdk_error_base(type(error))(
error.message, response=error.response, body=error.body
)
return sdk_error.with_traceback(error.__traceback__)
CLIENT_LOGGER_NAME = "composio_client"
"""Name of the logger the generated ``composio_client`` package writes to."""
class _ClientLogForwarder(logging.Handler):
"""Forward ``composio_client`` records into the SDK logger.
The generated client logs request/response lifecycle through
``logging.getLogger("composio_client")``. Routing those records through
the SDK's :class:`_VerbosityWrapper` keeps one destination for SDK users
and applies the same credential redaction and line truncation the SDK's
own records get. The client's INFO records (per-request lifecycle) are
forwarded as DEBUG.
"""
def __init__(self, wrapper: _VerbosityWrapper) -> None:
super().__init__()
self.wrapper = wrapper
def emit(self, record: logging.LogRecord) -> None:
try:
message = record.getMessage()
if record.levelno >= logging.ERROR:
self.wrapper.error(message, exc_info=record.exc_info)
elif record.levelno >= logging.WARNING:
self.wrapper.warning(message, exc_info=record.exc_info)
else:
self.wrapper.debug(message, exc_info=record.exc_info)
except Exception: # noqa: BLE001 - logging must never fail the call
self.handleError(record)
def _install_client_log_forwarder(wrapper: _VerbosityWrapper) -> logging.Logger:
"""Attach a single forwarder for ``wrapper`` to the client logger.
Idempotent: a forwarder already bound to ``wrapper`` is kept; forwarders
bound to another wrapper (an earlier SDK instance with a different logger)
are replaced so records are delivered once, to the most recent logger.
"""
client_logger = logging.getLogger(CLIENT_LOGGER_NAME)
installed = False
for handler in list(client_logger.handlers):
if not isinstance(handler, _ClientLogForwarder):
continue
if handler.wrapper is wrapper and not installed:
installed = True
continue
client_logger.removeHandler(handler)
if not installed:
client_logger.addHandler(_ClientLogForwarder(wrapper))
# The client's INFO records are forwarded as DEBUG, so only let the client
# produce them when the SDK logger is at DEBUG.
level = wrapper.logger.getEffectiveLevel()
client_logger.setLevel(
level if level <= logging.DEBUG else max(level, logging.WARNING)
)
client_logger.propagate = False
return client_logger
def _get_python_implementation() -> str:
"""
Get the Python implementation name.
Returns:
String identifier for Python implementation (CPYTHON, PYPY, JYTHON, IRONPYTHON, etc.)
"""
impl = platform.python_implementation().upper()
return impl
def _detect_runtime_environment() -> str:
"""
Detect the runtime environment where the code is executing.
Returns a string identifier for the environment.
"""
# Check for Google Colab
try:
import google.colab # type: ignore # noqa: F401
return "GOOGLE_COLAB"
except ImportError:
pass
# Check for Jupyter/IPython
try:
shell = get_ipython().__class__.__name__ # type: ignore # noqa: F821
if shell == "ZMQInteractiveShell":
return "JUPYTER_NOTEBOOK"
elif shell == "TerminalInteractiveShell":
return "IPYTHON"
except NameError:
pass
# Check for AWS Lambda
if os.environ.get("AWS_LAMBDA_FUNCTION_NAME"):
return "AWS_LAMBDA"
# Check for Google Cloud Functions
if os.environ.get("FUNCTION_NAME") or os.environ.get("K_SERVICE"):
return "GOOGLE_CLOUD_FUNCTION"
# Check for Azure Functions
if os.environ.get("FUNCTIONS_WORKER_RUNTIME"):
return "AZURE_FUNCTION"
# Check for Kaggle
if os.environ.get("KAGGLE_KERNEL_RUN_TYPE"):
return "KAGGLE"
# Check for Replit
if os.environ.get("REPL_ID") or os.environ.get("REPLIT_DB_URL"):
return "REPLIT"
# Check for GitHub Actions
if os.environ.get("GITHUB_ACTIONS"):
return "GITHUB_ACTIONS"
# Check for GitLab CI
if os.environ.get("GITLAB_CI"):
return "GITLAB_CI"
# Check for CircleCI
if os.environ.get("CIRCLECI"):
return "CIRCLECI"
# Check for Jenkins
if os.environ.get("JENKINS_HOME"):
return "JENKINS"
# Check for Docker
if os.path.exists("/.dockerenv") or os.path.exists("/run/.containerenv"):
return "DOCKER"
# Check if running in a container (generic)
try:
with open("/proc/1/cgroup", "r") as f:
if "docker" in f.read() or "containerd" in f.read():
return "CONTAINER"
except (FileNotFoundError, PermissionError):
pass
# Default to LOCAL for development environments
return "LOCAL"
class RequestContext(te.TypedDict):
id: te.NotRequired[t.Optional[str]]
provider: str
# TODO: Rename `Composio` to `HttpClient` in stainless generator
class HttpClient(BaseComposio, WithLogger):
"""
Wrapper around the auto-generated composio client.
"""
request_ctx: contextvars.ContextVar[RequestContext]
not_given = NOT_GIVEN
# Detect once at class initialization
_runtime_env: str = (
f"{_detect_runtime_environment()}_{_get_python_implementation()}"
)
def __init__(
self,
*,
provider: str,
api_key: t.Optional[str] = None,
user_api_key: t.Optional[str] = None,
org_api_key: t.Optional[str] = None,
environment: te.Union[NotGiven, APIEnvironment] = "production",
base_url: t.Optional[t.Union[str, URL, NotGiven]] = NOT_GIVEN,
timeout: t.Optional[t.Union[float, Timeout, NotGiven]] = NOT_GIVEN,
max_retries: int = DEFAULT_MAX_RETRIES,
default_headers: t.Optional[t.Mapping[str, str]] = None,
default_query: t.Optional[t.Mapping[str, object]] = None,
http_client: t.Optional[Client] = None,
logger: t.Optional[logging.Logger] = None,
logging_level: t.Optional[LogLevel] = None,
_strict_response_validation: bool = False,
) -> None:
"""
Initialize the client.
:param provider: The provider to use for the client.
:param logger: Logger that receives SDK and ``composio_client`` records.
:param logging_level: Level applied to the SDK and ``composio_client`` loggers.
:param api_key: The API key to use for the client.
:param user_api_key: User API key, sent only on operations that require it.
:param org_api_key: Organization API key, sent only on operations that require it.
:param environment: The environment to use for the client.
:param base_url: The base URL to use for the client.
:param timeout: The timeout to use for the client.
:param max_retries: The maximum number of retries to use for the client.
:param default_headers: The default headers to use for the client.
:param default_query: The default query parameters to use for the client.
:param http_client: The HTTP client to use for the client.
"""
WithLogger.__init__(self, logger=logger, logging_level=logging_level)
BaseComposio.__init__(
self,
api_key=api_key,
user_api_key=user_api_key,
org_api_key=org_api_key,
environment=environment,
base_url=base_url,
timeout=timeout,
max_retries=max_retries,
default_headers=default_headers,
default_query=default_query,
http_client=http_client,
_strict_response_validation=_strict_response_validation,
)
_install_client_log_forwarder(self._logger)
self.provider = provider
self.request_ctx = contextvars.ContextVar[RequestContext](
"request_ctx",
default={
"id": None,
"provider": provider,
},
)
# Lazily-built sibling client with retries disabled; see `without_retries`.
self._without_retries: t.Optional[te.Self] = None
def copy( # type: ignore[override]
self,
*,
_extra_kwargs: t.Mapping[str, t.Any] = {},
**kwargs: t.Any,
) -> te.Self:
"""
Clone the client, re-injecting the required ``provider`` keyword.
The Stainless-generated ``copy`` rebuilds the client via
``self.__class__(...)`` without passing ``provider``, which this subclass
requires — so the inherited ``copy``/``with_options`` raise ``TypeError``.
Threading ``provider`` through ``_extra_kwargs`` makes them work again
(e.g. ``with_options(max_retries=0)``).
"""
return super().copy( # type: ignore[misc]
_extra_kwargs={
"provider": self.provider,
# The generated `copy` does not re-pass `_strict_response_validation`,
# so without this the clone would silently fall back to the default
# (False) even when the original had it enabled — keeping the sibling
# a faithful copy that differs from the parent only in `max_retries`.
"_strict_response_validation": self._strict_response_validation,
# Share the parent's logger; otherwise constructing the clone
# rebinds the process-wide client log forwarder to the default
# `composio` logger.
"logger": self._logger.logger,
**_extra_kwargs,
},
**kwargs,
)
# Re-alias `with_options` to this override. The base class binds
# `with_options = copy` at class-definition time, so without this it would
# still resolve to the base `copy` and miss the `provider` re-injection.
with_options = copy
@property
def without_retries(self) -> te.Self:
"""
A cached sibling client that never retries requests.
Used for non-idempotent writes (``tools.execute`` / ``tools.proxy``),
where a silent retry after a read timeout can duplicate a side effect
(e.g. send an email twice). Reads keep the default retry behaviour.
Scope: only ``tools.execute`` / ``tools.proxy`` route through this today.
Other non-idempotent writes (``auth_configs.create`` / ``update`` /
``delete``, ``mcp.update`` / ``delete``, ``connected_accounts.delete`` /
``refresh``, ``link.create``) keep the default retries — most are
naturally idempotent on retry, and the durable fix is backend-honoured
idempotency keys.
The sibling is cached rather than rebuilt per call so a fresh client is
not constructed on every execute/proxy (the hottest path); its options
never change, so one per client suffices.
"""
if self._without_retries is None:
self._without_retries = self.with_options(max_retries=0)
return self._without_retries
def _decode(self, response: Response) -> t.Any:
"""
Raise status errors that are also ``ComposioError``s; see
``_with_sdk_error_base``.
"""
try:
return super()._decode(response)
except APIStatusError as error:
raise _as_sdk_error(error) from None
def _process_response(
self, response: Response, cast_to: t.Optional[t.Type[t.Any]]
) -> t.Any:
"""
Raise status errors that are also ``ComposioError``s; see
``_with_sdk_error_base``.
"""
try:
return super()._process_response(response, cast_to)
except APIStatusError as error:
raise _as_sdk_error(error) from None
def _prepare_request(self, request: Request) -> None:
"""
Request interceptor to inject request id, provider, and SDK version.
"""
ctx = self.request_ctx.get()
request.headers["x-request-id"] = ctx.get("id") or uuid4().hex
request.headers["x-framework"] = ctx["provider"]
request.headers["x-source"] = "PYTHON_SDK"
request.headers["x-runtime"] = HttpClient._runtime_env
try:
request.headers["x-sdk-version"] = version("composio")
except Exception:
request.headers["x-sdk-version"] = "unknown"