242 lines
8.9 KiB
Python
242 lines
8.9 KiB
Python
|
|
"""Tests for unified error handling infrastructure (B.13).
|
||
|
|
|
||
|
|
Covers:
|
||
|
|
- ApiError with code/category/retryable → correct response format
|
||
|
|
- HTTPException → unified error format
|
||
|
|
- Unhandled Exception → 500 with trace_id
|
||
|
|
- All defined error codes have correct fields
|
||
|
|
- classify_exception: transient/permanent/partial
|
||
|
|
"""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import pytest
|
||
|
|
from fastapi import HTTPException
|
||
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
|
|
||
|
|
from app.core.error_codes import (
|
||
|
|
ApiError,
|
||
|
|
ErrorCategory,
|
||
|
|
ERROR_CODES,
|
||
|
|
build_error_response,
|
||
|
|
classify_exception,
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
class TestApiError:
|
||
|
|
"""Tests for ApiError exception class."""
|
||
|
|
|
||
|
|
def test_api_error_basic(self):
|
||
|
|
err = ApiError("not_found")
|
||
|
|
assert err.code == "not_found"
|
||
|
|
assert err.status == 404
|
||
|
|
assert err.detail == "Resource not found"
|
||
|
|
assert err.field is None
|
||
|
|
assert err.category == ErrorCategory.PERMANENT
|
||
|
|
assert err.retryable is False
|
||
|
|
|
||
|
|
def test_api_error_with_custom_detail(self):
|
||
|
|
err = ApiError("validation_error", detail="Email is required", field="email")
|
||
|
|
assert err.code == "validation_error"
|
||
|
|
assert err.detail == "Email is required"
|
||
|
|
assert err.field == "email"
|
||
|
|
assert err.status == 422
|
||
|
|
assert err.category == ErrorCategory.PERMANENT
|
||
|
|
assert err.retryable is False
|
||
|
|
|
||
|
|
def test_api_error_transient(self):
|
||
|
|
err = ApiError("rate_limited")
|
||
|
|
assert err.category == ErrorCategory.TRANSIENT
|
||
|
|
assert err.retryable is True
|
||
|
|
assert err.status == 429
|
||
|
|
|
||
|
|
def test_api_error_with_explicit_category(self):
|
||
|
|
err = ApiError("internal_error", category=ErrorCategory.PERMANENT, retryable=False)
|
||
|
|
assert err.category == ErrorCategory.PERMANENT
|
||
|
|
assert err.retryable is False
|
||
|
|
|
||
|
|
def test_api_error_to_response(self):
|
||
|
|
err = ApiError("not_found", detail="Contact not found", field="id")
|
||
|
|
resp = err.to_response(trace_id="abc123")
|
||
|
|
assert resp["code"] == "not_found"
|
||
|
|
assert resp["detail"] == "Contact not found"
|
||
|
|
assert resp["field"] == "id"
|
||
|
|
assert resp["trace_id"] == "abc123"
|
||
|
|
assert resp["retryable"] is False
|
||
|
|
assert resp["category"] == "permanent"
|
||
|
|
|
||
|
|
def test_api_error_to_response_no_trace_id(self):
|
||
|
|
err = ApiError("rate_limited")
|
||
|
|
resp = err.to_response()
|
||
|
|
assert resp["trace_id"] is None
|
||
|
|
assert resp["retryable"] is True
|
||
|
|
assert resp["category"] == "transient"
|
||
|
|
|
||
|
|
|
||
|
|
class TestErrorCodes:
|
||
|
|
"""Tests for all defined error codes."""
|
||
|
|
|
||
|
|
def test_all_original_codes_exist(self):
|
||
|
|
expected = {
|
||
|
|
"not_found", "permission_denied", "validation_error",
|
||
|
|
"rate_limited", "internal_error", "service_unavailable",
|
||
|
|
}
|
||
|
|
assert expected.issubset(ERROR_CODES.keys())
|
||
|
|
|
||
|
|
def test_new_codes_exist(self):
|
||
|
|
new_codes = {
|
||
|
|
"forbidden", "conflict", "unprocessable",
|
||
|
|
"not_implemented", "service_timeout", "bad_gateway",
|
||
|
|
}
|
||
|
|
assert new_codes.issubset(ERROR_CODES.keys())
|
||
|
|
|
||
|
|
def test_all_codes_have_required_fields(self):
|
||
|
|
for code, meta in ERROR_CODES.items():
|
||
|
|
assert "status" in meta, f"{code} missing status"
|
||
|
|
assert "message" in meta, f"{code} missing message"
|
||
|
|
assert "category" in meta, f"{code} missing category"
|
||
|
|
assert "retryable" in meta, f"{code} missing retryable"
|
||
|
|
assert isinstance(meta["category"], ErrorCategory)
|
||
|
|
assert isinstance(meta["retryable"], bool)
|
||
|
|
assert isinstance(meta["status"], int)
|
||
|
|
assert 100 <= meta["status"] < 600
|
||
|
|
|
||
|
|
def test_forbidden_code(self):
|
||
|
|
meta = ERROR_CODES["forbidden"]
|
||
|
|
assert meta["status"] == 403
|
||
|
|
assert meta["category"] == ErrorCategory.PERMANENT
|
||
|
|
assert meta["retryable"] is False
|
||
|
|
|
||
|
|
def test_conflict_code(self):
|
||
|
|
meta = ERROR_CODES["conflict"]
|
||
|
|
assert meta["status"] == 409
|
||
|
|
assert meta["category"] == ErrorCategory.PERMANENT
|
||
|
|
|
||
|
|
def test_service_timeout_code(self):
|
||
|
|
meta = ERROR_CODES["service_timeout"]
|
||
|
|
assert meta["status"] == 504
|
||
|
|
assert meta["category"] == ErrorCategory.TRANSIENT
|
||
|
|
assert meta["retryable"] is True
|
||
|
|
|
||
|
|
def test_bad_gateway_code(self):
|
||
|
|
meta = ERROR_CODES["bad_gateway"]
|
||
|
|
assert meta["status"] == 502
|
||
|
|
assert meta["category"] == ErrorCategory.TRANSIENT
|
||
|
|
assert meta["retryable"] is True
|
||
|
|
|
||
|
|
|
||
|
|
class TestClassifyException:
|
||
|
|
"""Tests for classify_exception helper."""
|
||
|
|
|
||
|
|
def test_classify_timeout_transient(self):
|
||
|
|
exc = TimeoutError("Request timed out")
|
||
|
|
assert classify_exception(exc) == ErrorCategory.TRANSIENT
|
||
|
|
|
||
|
|
def test_classify_rate_limit_transient(self):
|
||
|
|
exc = Exception("rate limit exceeded")
|
||
|
|
assert classify_exception(exc) == ErrorCategory.TRANSIENT
|
||
|
|
|
||
|
|
def test_classify_connection_error_transient(self):
|
||
|
|
exc = ConnectionError("connection reset")
|
||
|
|
assert classify_exception(exc) == ErrorCategory.TRANSIENT
|
||
|
|
|
||
|
|
def test_classify_auth_permanent(self):
|
||
|
|
exc = Exception("authentication failed")
|
||
|
|
assert classify_exception(exc) == ErrorCategory.PERMANENT
|
||
|
|
|
||
|
|
def test_classify_validation_permanent(self):
|
||
|
|
exc = Exception("validation error: invalid input")
|
||
|
|
assert classify_exception(exc) == ErrorCategory.PERMANENT
|
||
|
|
|
||
|
|
def test_classify_permission_permanent(self):
|
||
|
|
exc = Exception("forbidden: permission denied")
|
||
|
|
assert classify_exception(exc) == ErrorCategory.PERMANENT
|
||
|
|
|
||
|
|
def test_classify_not_found_permanent(self):
|
||
|
|
exc = Exception("not found")
|
||
|
|
assert classify_exception(exc) == ErrorCategory.PERMANENT
|
||
|
|
|
||
|
|
def test_classify_partial_batch(self):
|
||
|
|
exc = Exception("batch operation partially failed")
|
||
|
|
assert classify_exception(exc) == ErrorCategory.PARTIAL
|
||
|
|
|
||
|
|
def test_classify_partial_bulk(self):
|
||
|
|
exc = Exception("bulk import: some failed")
|
||
|
|
assert classify_exception(exc) == ErrorCategory.PARTIAL
|
||
|
|
|
||
|
|
def test_classify_api_error_uses_own_category(self):
|
||
|
|
err = ApiError("rate_limited")
|
||
|
|
assert classify_exception(err) == ErrorCategory.TRANSIENT
|
||
|
|
|
||
|
|
def test_classify_api_error_permanent(self):
|
||
|
|
err = ApiError("not_found")
|
||
|
|
assert classify_exception(err) == ErrorCategory.PERMANENT
|
||
|
|
|
||
|
|
def test_classify_unknown_defaults_transient(self):
|
||
|
|
exc = Exception("some unknown error")
|
||
|
|
assert classify_exception(exc) == ErrorCategory.TRANSIENT
|
||
|
|
|
||
|
|
|
||
|
|
class TestBuildErrorResponse:
|
||
|
|
"""Tests for build_error_response helper."""
|
||
|
|
|
||
|
|
def test_build_error_response_basic(self):
|
||
|
|
resp = build_error_response("not_found", trace_id="xyz789")
|
||
|
|
assert resp["code"] == "not_found"
|
||
|
|
assert resp["detail"] == "Resource not found"
|
||
|
|
assert resp["trace_id"] == "xyz789"
|
||
|
|
assert resp["retryable"] is False
|
||
|
|
assert resp["category"] == "permanent"
|
||
|
|
|
||
|
|
def test_build_error_response_with_detail(self):
|
||
|
|
resp = build_error_response("validation_error", detail="Bad input")
|
||
|
|
assert resp["detail"] == "Bad input"
|
||
|
|
assert resp["category"] == "permanent"
|
||
|
|
|
||
|
|
|
||
|
|
class TestErrorResponsesViaAPI:
|
||
|
|
"""Integration tests for error response format via HTTP client."""
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_api_error_response_format(self, client):
|
||
|
|
"""ApiError should return unified format with all fields."""
|
||
|
|
# Use the /api/v1/errors/trigger endpoint if it exists
|
||
|
|
# Otherwise we test via a known 404 endpoint
|
||
|
|
response = await client.get("/api/v1/nonexistent-endpoint")
|
||
|
|
assert response.status_code in (404, 405)
|
||
|
|
body = response.json()
|
||
|
|
# Should have unified format fields
|
||
|
|
assert "code" in body
|
||
|
|
assert "detail" in body
|
||
|
|
assert "trace_id" in body
|
||
|
|
assert "retryable" in body
|
||
|
|
assert "category" in body
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_trace_id_in_response_header(self, client):
|
||
|
|
"""X-Trace-Id header should be present in responses."""
|
||
|
|
response = await client.get("/api/v1/health")
|
||
|
|
assert "x-trace-id" in response.headers
|
||
|
|
trace_id = response.headers["x-trace-id"]
|
||
|
|
assert len(trace_id) == 8 # 8-char hex
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_health_endpoint_success(self, client):
|
||
|
|
"""Health endpoint should return 200 and trace_id header."""
|
||
|
|
response = await client.get("/api/v1/health")
|
||
|
|
assert response.status_code == 200
|
||
|
|
assert "x-trace-id" in response.headers
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_404_has_unified_format(self, client):
|
||
|
|
"""404 responses should have unified error format."""
|
||
|
|
response = await client.get("/api/v1/contacts/00000000-0000-0000-0000-000000000000")
|
||
|
|
# Could be 401 (no auth) or 404 — both should have unified format
|
||
|
|
assert response.status_code in (401, 403, 404)
|
||
|
|
body = response.json()
|
||
|
|
assert "code" in body
|
||
|
|
assert "detail" in body
|
||
|
|
assert "trace_id" in body
|
||
|
|
assert "retryable" in body
|
||
|
|
assert "category" in body
|