from typing import Any import pytest from fastapi import FastAPI, HTTPException from fastapi.exceptions import RequestValidationError from pydantic import BaseModel from pydantic_core import ErrorDetails from pytest_mock import MockerFixture from starlette import status from starlette.exceptions import HTTPException as StarletteHTTPException from starlette.testclient import TestClient from fansifter_common.api.error_handlers import register_defaults from fansifter_common.auth.exceptions import NotAuthenticated, PermissionDenied from fansifter_common.exceptions import FansifterError class SimpleModel(BaseModel): field1: str field2: int def raise_exception() -> None: pass def validation_exc_route(data: SimpleModel) -> Any: return data def exc_route() -> Any: raise_exception() @pytest.fixture def client() -> TestClient: app = FastAPI(debug=False) app.add_api_route("/exc", endpoint=exc_route) app.add_api_route( "/exc-validation", methods=["POST"], endpoint=validation_exc_route, ) register_defaults(app) return TestClient(app, raise_server_exceptions=False) class CustomError(FansifterError): code = "custom_error" message = "Custom error" additional_properties = ( ("param_one", str), ("param_two", int), ) def __init__(self, param_one: str, param_two: int) -> None: super().__init__() self.param_one = param_one self.param_two = param_two @pytest.mark.parametrize( "exc, status_code, content", [ ( HTTPException(status_code=400, detail="Bad request"), status.HTTP_400_BAD_REQUEST, { "code": "bad_request", "message": "Bad request", }, ), ( StarletteHTTPException(status_code=400, detail="Bad request"), status.HTTP_400_BAD_REQUEST, { "code": "bad_request", "message": "Bad request", }, ), ( RequestValidationError(errors=[]), status.HTTP_422_UNPROCESSABLE_ENTITY, { "code": "invalid_input", "message": "Invalid input", "fieldErrors": {}, }, ), ( RequestValidationError( errors=[ ErrorDetails( loc=("body", "name"), msg="Field required", input="", type="required", ) ] ), status.HTTP_422_UNPROCESSABLE_ENTITY, { "code": "invalid_input", "message": "Invalid input", "fieldErrors": { "name": { "code": "required", "message": "Field required", }, }, }, ), ( FansifterError(), status.HTTP_500_INTERNAL_SERVER_ERROR, {"code": "fansifter_error", "message": "Fansifter error"}, ), ( FansifterError(message="Some error"), status.HTTP_500_INTERNAL_SERVER_ERROR, {"code": "fansifter_error", "message": "Some error"}, ), ( CustomError(param_one="one", param_two=2), status.HTTP_500_INTERNAL_SERVER_ERROR, { "code": "custom_error", "message": "Custom error", "paramOne": "one", "paramTwo": 2, }, ), ( NotAuthenticated, status.HTTP_401_UNAUTHORIZED, { "code": "authorization_error", "message": "Unauthorized", }, ), ( NotAuthenticated("missing header"), status.HTTP_401_UNAUTHORIZED, { "code": "authorization_error", "message": "missing header", }, ), ( PermissionDenied("not allowed"), status.HTTP_403_FORBIDDEN, { "code": "permission_denied", "message": "not allowed", }, ), ], ) def test_error_handler( exc: Exception | type[Exception], status_code: int, content: Any, client: TestClient, mocker: MockerFixture, ) -> None: mocker.patch(f"{__name__}.raise_exception", side_effect=exc) response = client.get("/exc") assert response.status_code == status_code assert response.json() == content def test_request_validation_error_handler(client: TestClient) -> None: response = client.post("/exc-validation", json={"field2": "invalid"}) assert response.status_code == status.HTTP_422_UNPROCESSABLE_ENTITY assert response.json() == { "code": "invalid_input", "message": "Invalid input", "fieldErrors": { "field1": { "code": "missing", "message": "Field required", }, "field2": { "code": "int_parsing", "message": ( "Input should be a valid integer, " "unable to parse string as an integer" ), }, }, }