"""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