fix: api_error_handler silently drops return values from async methods (#5540)

This commit is contained in:
Hrushikesh Yadav
2026-06-17 14:46:08 +05:30
committed by GitHub
parent 0d66d3d127
commit 7ed2faab84
2 changed files with 174 additions and 87 deletions
+99 -87
View File
@@ -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
+75
View File
@@ -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"