diff --git a/test/api/test_api.py b/test/api/test_api.py index 83e2b23..d381b5e 100644 --- a/test/api/test_api.py +++ b/test/api/test_api.py @@ -581,3 +581,41 @@ def test_no_param_plugin_still_accepts_a_bodyless_post(): assert resp.status_code == 200 assert InvokeResponse.model_validate(resp.json()).output["received"] == "ok" + + +class _PrecheckFailure(Exception): + status_code = 403 + failure_category = "AUTH_PERMISSION_DENIED" + + +def _failing_precheck() -> None: + raise _PrecheckFailure("credential rejected") + + +def _passing_precheck() -> None: + return None + + +def test_precheck_reports_failure_category_from_raised_error(): + client = TestClient( + wrap_in_fastapi(func=_no_params, plugin_id="mock_plugin", precheck_func=_failing_precheck) + ) + + resp = client.get("/precheck") + + body = resp.json() + assert body["status_code"] == 403 + assert body["failure_category"] == "AUTH_PERMISSION_DENIED" + assert "credential rejected" in body["status_code_text"] + + +def test_precheck_success_has_no_failure_category(): + client = TestClient( + wrap_in_fastapi(func=_no_params, plugin_id="mock_plugin", precheck_func=_passing_precheck) + ) + + resp = client.get("/precheck") + + body = resp.json() + assert body["status_code"] == 200 + assert body["failure_category"] is None diff --git a/unstructured_platform_plugins/etl_uvicorn/api_generator.py b/unstructured_platform_plugins/etl_uvicorn/api_generator.py index 4b7d8ce..168a3d3 100644 --- a/unstructured_platform_plugins/etl_uvicorn/api_generator.py +++ b/unstructured_platform_plugins/etl_uvicorn/api_generator.py @@ -146,6 +146,7 @@ class InvokeResponse(BaseModel): file_data: Optional[FileDataType] = None filedata_meta: Optional[filedata_meta_model] = None status_code_text: Optional[str] = None + failure_category: Optional[str] = None output: Optional[response_type] = None message_channels: MessageChannels = Field(default_factory=MessageChannels) @@ -200,6 +201,7 @@ async def _stream_response(): status_code=getattr(e, "status_code", None) or status.HTTP_500_INTERNAL_SERVER_ERROR, status_code_text=f"[{e.__class__.__name__}] {e}", + failure_category=getattr(e, "failure_category", None), ).model_dump_json() + "\n" ) @@ -227,6 +229,7 @@ async def _stream_response(): status_code_text=json.dumps(exc.detail) if isinstance(exc.detail, dict) else exc.detail, + failure_category=getattr(exc, "failure_category", None), file_data=request_dict.get("file_data", None), ) except UnstructuredIngestError as exc: @@ -240,6 +243,7 @@ async def _stream_response(): filedata_meta=filedata_meta_model.model_validate(filedata_meta.model_dump()), status_code=exc.status_code or status.HTTP_500_INTERNAL_SERVER_ERROR, status_code_text=str(exc), + failure_category=getattr(exc, "failure_category", None), file_data=request_dict.get("file_data", None), ) except Exception as invoke_error: @@ -251,6 +255,7 @@ async def _stream_response(): status_code=getattr(invoke_error, "status_code", None) or status.HTTP_500_INTERNAL_SERVER_ERROR, status_code_text=f"[{invoke_error.__class__.__name__}] {invoke_error}", + failure_category=getattr(invoke_error, "failure_category", None), file_data=request_dict.get("file_data", None), ) @@ -318,6 +323,7 @@ class InvokePrecheckResponse(BaseModel): usage: list[UsageData] status_code: int status_code_text: Optional[str] = None + failure_category: Optional[str] = None @fastapi_app.get("/schema") async def get_schema() -> SchemaOutputResponse: @@ -332,6 +338,7 @@ async def run_precheck() -> InvokePrecheckResponse: return InvokePrecheckResponse( status_code=fn_response.status_code, status_code_text=fn_response.status_code_text, + failure_category=fn_response.failure_category, usage=fn_response.usage, ) else: