diff --git a/mem0/client/utils.py b/mem0/client/utils.py index 06a1c0ef2..57fe42707 100644 --- a/mem0/client/utils.py +++ b/mem0/client/utils.py @@ -1,5 +1,8 @@ +import inspect import json import logging +from functools import wraps + import httpx from mem0.exceptions import ( @@ -12,7 +15,7 @@ logger = logging.getLogger(__name__) class APIError(Exception): """Exception raised for errors in the API. - + Deprecated: Use specific exception classes from mem0.exceptions instead. This class is maintained for backward compatibility. """ @@ -20,96 +23,105 @@ class APIError(Exception): pass +def _handle_http_error(e): + logger.error(f"HTTP error occurred: {e}") + + response_text = "" + error_details = {} + debug_info = { + "status_code": e.response.status_code, + "url": str(e.request.url), + "method": e.request.method, + } + + try: + response_text = e.response.text + if e.response.headers.get("content-type", "").startswith("application/json"): + error_data = json.loads(response_text) + if isinstance(error_data, dict): + error_details = error_data + response_text = error_data.get("detail", response_text) + except (json.JSONDecodeError, AttributeError): + pass + + if e.response.status_code == 429: + retry_after = e.response.headers.get("Retry-After") + if retry_after: + try: + debug_info["retry_after"] = int(retry_after) + except ValueError: + pass + + for header in ["X-RateLimit-Limit", "X-RateLimit-Remaining", "X-RateLimit-Reset"]: + value = e.response.headers.get(header) + if value: + debug_info[header.lower().replace("-", "_")] = value + + raise create_exception_from_response( + status_code=e.response.status_code, + response_text=response_text, + details=error_details, + debug_info=debug_info, + ) + + +def _handle_request_error(e): + logger.error(f"Request error occurred: {e}") + + if isinstance(e, httpx.TimeoutException): + raise NetworkError( + message=f"Request timed out: {str(e)}", + error_code="NET_TIMEOUT", + suggestion="Please check your internet connection and try again", + debug_info={"error_type": "timeout", "original_error": str(e)}, + ) + elif isinstance(e, httpx.ConnectError): + raise NetworkError( + message=f"Connection failed: {str(e)}", + error_code="NET_CONNECT", + suggestion="Please check your internet connection and try again", + debug_info={"error_type": "connection", "original_error": str(e)}, + ) + else: + raise NetworkError( + message=f"Network request failed: {str(e)}", + error_code="NET_GENERIC", + suggestion="Please check your internet connection and try again", + debug_info={"error_type": "request", "original_error": str(e)}, + ) + + def api_error_handler(func): """Decorator to handle API errors consistently. - + This decorator catches HTTP and request errors and converts them to appropriate structured exception classes with detailed error information. - - The decorator analyzes HTTP status codes and response content to create - the most specific exception type with helpful error messages, suggestions, - and debug information. + + Supports both sync and async functions. """ - from functools import wraps - - @wraps(func) - def wrapper(*args, **kwargs): - try: - return func(*args, **kwargs) - except httpx.HTTPStatusError as e: - logger.error(f"HTTP error occurred: {e}") - - # Extract error details from response - response_text = "" - error_details = {} - debug_info = { - "status_code": e.response.status_code, - "url": str(e.request.url), - "method": e.request.method, - } - + if inspect.iscoroutinefunction(func): + @wraps(func) + async def async_wrapper(*args, **kwargs): try: - response_text = e.response.text - # Try to parse JSON response for additional error details - if e.response.headers.get("content-type", "").startswith("application/json"): - error_data = json.loads(response_text) - if isinstance(error_data, dict): - error_details = error_data - response_text = error_data.get("detail", response_text) - except (json.JSONDecodeError, AttributeError): - # Fallback to plain text response - pass - - # Add rate limit information if available - if e.response.status_code == 429: - retry_after = e.response.headers.get("Retry-After") - if retry_after: - try: - debug_info["retry_after"] = int(retry_after) - except ValueError: - pass - - # Add rate limit headers if available - for header in ["X-RateLimit-Limit", "X-RateLimit-Remaining", "X-RateLimit-Reset"]: - value = e.response.headers.get(header) - if value: - debug_info[header.lower().replace("-", "_")] = value - - # Create specific exception based on status code - exception = create_exception_from_response( - status_code=e.response.status_code, - response_text=response_text, - details=error_details, - debug_info=debug_info, - ) - - raise exception - - except httpx.RequestError as e: - logger.error(f"Request error occurred: {e}") - - # Determine the appropriate exception type based on error type - if isinstance(e, httpx.TimeoutException): - raise NetworkError( - message=f"Request timed out: {str(e)}", - error_code="NET_TIMEOUT", - suggestion="Please check your internet connection and try again", - debug_info={"error_type": "timeout", "original_error": str(e)}, - ) - elif isinstance(e, httpx.ConnectError): - raise NetworkError( - message=f"Connection failed: {str(e)}", - error_code="NET_CONNECT", - suggestion="Please check your internet connection and try again", - debug_info={"error_type": "connection", "original_error": str(e)}, - ) - else: - # Generic network error for other request errors - raise NetworkError( - message=f"Network request failed: {str(e)}", - error_code="NET_GENERIC", - suggestion="Please check your internet connection and try again", - debug_info={"error_type": "request", "original_error": str(e)}, - ) + return await func(*args, **kwargs) + except httpx.HTTPStatusError as e: + _handle_http_error(e) + raise + except httpx.RequestError as e: + _handle_request_error(e) + raise - return wrapper + return async_wrapper + else: + @wraps(func) + def wrapper(*args, **kwargs): + try: + return func(*args, **kwargs) + except httpx.HTTPStatusError as e: + _handle_http_error(e) + raise + except httpx.RequestError as e: + _handle_request_error(e) + raise + + return wrapper diff --git a/tests/test_client_utils.py b/tests/test_client_utils.py new file mode 100644 index 000000000..d5f4771c5 --- /dev/null +++ b/tests/test_client_utils.py @@ -0,0 +1,75 @@ +import inspect + +import httpx +import pytest + +from mem0.client.utils import api_error_handler +from mem0.exceptions import AuthenticationError, NetworkError, RateLimitError + + +def test_sync_returns_value(): + @api_error_handler + def sync_fn(): + return {"result": "ok"} + + assert sync_fn() == {"result": "ok"} + + +@pytest.mark.asyncio +async def test_async_returns_value(): + @api_error_handler + async def async_fn(): + return {"result": "ok"} + + assert await async_fn() == {"result": "ok"} + + +def test_async_decorated_preserves_coroutine_flag(): + @api_error_handler + async def async_fn(): + return True + + assert inspect.iscoroutinefunction(async_fn) + + +def test_sync_decorated_is_not_coroutine(): + @api_error_handler + def sync_fn(): + return True + + assert not inspect.iscoroutinefunction(sync_fn) + + +def test_sync_http_error_raises_structured_exception(): + @api_error_handler + def sync_fn(): + request = httpx.Request("GET", "https://api.mem0.ai/v1/memories") + response = httpx.Response(401, request=request, text="Unauthorized") + raise httpx.HTTPStatusError("401", request=request, response=response) + + with pytest.raises(AuthenticationError): + sync_fn() + + +@pytest.mark.asyncio +async def test_async_http_error_raises_structured_exception(): + @api_error_handler + async def async_fn(): + request = httpx.Request("GET", "https://api.mem0.ai/v1/memories") + response = httpx.Response(429, request=request, text="Rate limited") + raise httpx.HTTPStatusError("429", request=request, response=response) + + with pytest.raises(RateLimitError): + await async_fn() + + +@pytest.mark.asyncio +async def test_async_connect_error_raises_network_error(): + @api_error_handler + async def async_fn(): + request = httpx.Request("GET", "https://api.mem0.ai/v1/memories") + raise httpx.ConnectError("Connection refused", request=request) + + with pytest.raises(NetworkError) as exc_info: + await async_fn() + assert exc_info.value.error_code == "NET_CONNECT"