diff --git a/test/api/test_api.py b/test/api/test_api.py index 83e2b23..61a6c94 100644 --- a/test/api/test_api.py +++ b/test/api/test_api.py @@ -1,5 +1,7 @@ +import logging from pathlib import Path from typing import Any, Optional, Union +from unittest.mock import patch import pytest from fastapi.testclient import TestClient @@ -581,3 +583,248 @@ 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 + + +def test_precheck_ignores_non_string_failure_category(): + class _NonStringCategoryFailure(Exception): + status_code = 403 + failure_category = 403 + + def _enum_category_precheck() -> None: + raise _NonStringCategoryFailure("credential rejected") + + client = TestClient( + wrap_in_fastapi( + func=_no_params, plugin_id="mock_plugin", precheck_func=_enum_category_precheck + ) + ) + + resp = client.get("/precheck") + + body = resp.json() + assert body["status_code"] == 403 + assert body["failure_category"] is None + assert "credential rejected" in body["status_code_text"] + + +def test_invoke_reports_failure_category_from_raised_error(): + client = TestClient(wrap_in_fastapi(func=_failing_precheck, plugin_id="mock_plugin")) + + body = client.post("/invoke").json() + assert body["status_code"] == 403 + assert body["failure_category"] == "AUTH_PERMISSION_DENIED" + + +def test_invoke_sanitizes_raising_error_attributes(): + class _HostileError(Exception): + @property + def status_code(self) -> int: + raise RuntimeError("status_code exploded") + + @property + def failure_category(self) -> str: + raise RuntimeError("failure_category exploded") + + def _raising_func() -> None: + raise _HostileError("original message") + + client = TestClient(wrap_in_fastapi(func=_raising_func, plugin_id="mock_plugin")) + + resp = client.post("/invoke") + assert resp.status_code == 200 + body = resp.json() + assert body["status_code"] == 500 + assert body["failure_category"] is None + assert "original message" in body["status_code_text"] + + +class _UnrenderableError(Exception): + status_code = 403 + + def __str__(self) -> str: + raise RuntimeError("__str__ exploded") + + +def test_invoke_survives_error_whose_str_raises(): + def _raising_func() -> None: + raise _UnrenderableError() + + client = TestClient(wrap_in_fastapi(func=_raising_func, plugin_id="mock_plugin")) + + resp = client.post("/invoke") + assert resp.status_code == 200 + body = resp.json() + assert body["status_code"] == 403 + assert "" in body["status_code_text"] + + +def test_invoke_survives_error_with_hostile_class_access(): + class _HostileClassError(Exception): + status_code = 403 + + def __getattribute__(self, name: str): + if name == "__class__": + raise RuntimeError("__class__ exploded") + return super().__getattribute__(name) + + def _raising_func() -> None: + raise _HostileClassError("original message") + + client = TestClient(wrap_in_fastapi(func=_raising_func, plugin_id="mock_plugin")) + + # The handler logs the error with exc_info; formatting this exception's + # traceback outside the handler would raise, so keep the record out of + # the captured-log machinery. + with patch.object(logging.getLogger("uvicorn.error"), "error"): + resp = client.post("/invoke") + + assert resp.status_code == 200 + body = resp.json() + assert body["status_code"] == 403 + assert "_HostileClassError" in body["status_code_text"] + assert "original message" in body["status_code_text"] + + +def test_precheck_survives_error_whose_str_raises(): + def _unrenderable_precheck() -> None: + raise _UnrenderableError() + + client = TestClient( + wrap_in_fastapi( + func=_no_params, plugin_id="mock_plugin", precheck_func=_unrenderable_precheck + ) + ) + + resp = client.get("/precheck") + assert resp.status_code == 200 + body = resp.json() + assert body["status_code"] == 403 + assert "" in body["status_code_text"] + + +def test_invoke_ignores_non_integer_status_code(): + class _BadStatusError(Exception): + status_code = "not-a-code" + + def _raising_func() -> None: + raise _BadStatusError("boom") + + client = TestClient(wrap_in_fastapi(func=_raising_func, plugin_id="mock_plugin")) + + body = client.post("/invoke").json() + assert body["status_code"] == 500 + + +def test_invoke_serializes_non_string_http_exception_detail(): + from fastapi import HTTPException + + def _raising_func() -> None: + raise HTTPException(status_code=422, detail=["field a", "field b"]) + + client = TestClient(wrap_in_fastapi(func=_raising_func, plugin_id="mock_plugin")) + + body = client.post("/invoke").json() + assert body["status_code"] == 422 + assert body["status_code_text"] == '["field a", "field b"]' + + +def test_precheck_func_may_take_a_usage_list_parameter(): + def _usage_precheck(usage: list) -> None: + return None + + client = TestClient( + wrap_in_fastapi(func=_no_params, plugin_id="mock_plugin", precheck_func=_usage_precheck) + ) + + assert client.get("/precheck").json()["status_code"] == 200 + + +def test_precheck_func_with_non_list_usage_parameter_is_rejected(): + from unstructured_platform_plugins.etl_uvicorn.api_generator import EtlApiException + + def _bad_precheck(usage: int) -> None: + return None + + with pytest.raises(EtlApiException): + wrap_in_fastapi(func=_no_params, plugin_id="mock_plugin", precheck_func=_bad_precheck) + + +def test_invoke_clamps_out_of_range_status_code(): + class _ZeroStatusError(Exception): + status_code = 0 + + def _raising_func() -> None: + raise _ZeroStatusError("boom") + + client = TestClient(wrap_in_fastapi(func=_raising_func, plugin_id="mock_plugin")) + + assert client.post("/invoke").json()["status_code"] == 500 + + +def test_invoke_survives_ingest_error_with_raising_status_code(): + from unstructured_ingest.error import UnstructuredIngestError + + class _HostileIngestError(UnstructuredIngestError): + @property + def status_code(self) -> int: + raise RuntimeError("status_code exploded") + + def _raising_func() -> None: + raise _HostileIngestError("boom") + + client = TestClient(wrap_in_fastapi(func=_raising_func, plugin_id="mock_plugin")) + + resp = client.post("/invoke") + assert resp.status_code == 200 + assert resp.json()["status_code"] == 500 + + +def test_precheck_func_accepts_string_annotations(): + def _string_annotated_precheck(usage: "list") -> "None": + return None + + client = TestClient( + wrap_in_fastapi( + func=_no_params, plugin_id="mock_plugin", precheck_func=_string_annotated_precheck + ) + ) + + assert client.get("/precheck").json()["status_code"] == 200 diff --git a/unstructured_platform_plugins/etl_uvicorn/api_generator.py b/unstructured_platform_plugins/etl_uvicorn/api_generator.py index 4b7d8ce..5534260 100644 --- a/unstructured_platform_plugins/etl_uvicorn/api_generator.py +++ b/unstructured_platform_plugins/etl_uvicorn/api_generator.py @@ -3,8 +3,7 @@ import inspect import json import logging -from functools import partial -from typing import Any, Callable, Optional, Union +from typing import Any, Callable, Optional, Union, get_origin from fastapi import FastAPI, HTTPException, status from fastapi.responses import StreamingResponse @@ -62,21 +61,74 @@ def log_func_and_body(func: Callable, body: Optional[str] = None) -> None: logger.log(level=logger.level, msg=msg) +def _error_attr(error: BaseException, name: str) -> Any: + """Read an attribute off a raised error, treating a raising property as absent. + + Runs while an exception handler is building the sanitized response; an + attribute access that itself raises must not replace that response with a + raw 500. + """ + try: + return getattr(error, name, None) + except Exception: + return None + + +def _safe_str(value: object) -> str: + """str() on a plugin-supplied error can itself raise; never let that escape the handler. + + An escape replaces the sanitized envelope with a raw HTTP 500, which the + controller's preflight reads as fail-open — a plugin-reported failure would + silently become a proceed. + """ + try: + return str(value) + except Exception: + return "" + + +def failure_category_of(error: BaseException) -> Optional[str]: + """Return the error's failure_category only when it is a plain string.""" + category = _error_attr(error, "failure_category") + return category if isinstance(category, str) else None + + +def status_code_of(error: BaseException) -> int: + """Return the error's status_code only when it is an int in the HTTP range, else 500.""" + status_code = _error_attr(error, "status_code") + if ( + isinstance(status_code, int) + and not isinstance(status_code, bool) + and 100 <= status_code <= 599 + ): + return status_code + return status.HTTP_500_INTERNAL_SERVER_ERROR + + async def invoke_func(func: Callable, kwargs: Optional[dict[str, Any]] = None) -> Any: kwargs = kwargs or {} if inspect.iscoroutinefunction(func): return await func(**kwargs) - else: - return await asyncio.get_event_loop().run_in_executor(None, partial(func, **kwargs)) + # to_thread copies contextvars into the worker thread, so OpenTelemetry + # context (and any wide-event adoption inside func) survives the hop; + # run_in_executor does not. + return await asyncio.to_thread(func, **kwargs) def check_precheck_func(precheck_func: Callable): - sig = inspect.signature(precheck_func) - inputs = sig.parameters.values() + try: + # eval_str resolves postponed/string annotations ('list', 'None') + sig = inspect.signature(precheck_func, eval_str=True) + except (NameError, TypeError): + sig = inspect.signature(precheck_func) + inputs = list(sig.parameters.values()) outputs = sig.return_annotation if len(inputs) == 1: i = inputs[0] - if i.name != "usage" or i.annotation is list: + annotation_is_list = ( + i.annotation is sig.empty or i.annotation is list or get_origin(i.annotation) is list + ) + if i.name != "usage" or not annotation_is_list: raise ValueError("the only input available for precheck is usage which must be a list") if outputs not in [None, sig.empty]: raise ValueError(f"no output should exist for precheck function, found: {outputs}") @@ -135,6 +187,9 @@ def _wrap_in_fastapi( logger.debug(f"set static id response to: {plugin_id}") + if "usage" not in inspect.signature(func).parameters: + logger.warning("usage data not an expected parameter, omitting") + fastapi_app = FastAPI() response_type = get_output_sig(func) @@ -146,6 +201,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) @@ -161,13 +217,12 @@ async def wrap_fn(func: Callable, kwargs: Optional[dict[str, Any]] = None) -> Re filedata_meta = FileDataMeta() message_channels = MessageChannels() request_dict = kwargs if kwargs else {} - if "usage" in inspect.signature(func).parameters: + params = inspect.signature(func).parameters + if "usage" in params: request_dict["usage"] = usage - else: - logger.warning("usage data not an expected parameter, omitting") - if "message_channels" in inspect.signature(func).parameters: + if "message_channels" in params: request_dict["message_channels"] = message_channels - if "filedata_meta" in inspect.signature(func).parameters: + if "filedata_meta" in params: request_dict["filedata_meta"] = filedata_meta try: if inspect.isasyncgenfunction(func): @@ -189,7 +244,7 @@ async def _stream_response(): + "\n" ) except Exception as e: - logger.error(f"Failure streaming response: {e}", exc_info=True) + logger.error(f"Failure streaming response: {_safe_str(e)}", exc_info=True) yield ( InvokeResponse( usage=usage, @@ -197,9 +252,9 @@ async def _stream_response(): filedata_meta=filedata_meta_model.model_validate( filedata_meta.model_dump() ), - status_code=getattr(e, "status_code", None) - or status.HTTP_500_INTERNAL_SERVER_ERROR, - status_code_text=f"[{e.__class__.__name__}] {e}", + status_code=status_code_of(e), + status_code_text=f"[{type(e).__name__}] {_safe_str(e)}", + failure_category=failure_category_of(e), ).model_dump_json() + "\n" ) @@ -217,40 +272,44 @@ async def _stream_response(): ) except HTTPException as exc: logger.error( - f"HTTPException: {exc.detail} (status_code={exc.status_code})", exc_info=True + f"HTTPException: {_safe_str(exc.detail)} (status_code={exc.status_code})", + exc_info=True, ) return InvokeResponse( usage=usage, message_channels=message_channels, filedata_meta=filedata_meta_model.model_validate(filedata_meta.model_dump()), status_code=exc.status_code, - status_code_text=json.dumps(exc.detail) - if isinstance(exc.detail, dict) - else exc.detail, + status_code_text=exc.detail + if isinstance(exc.detail, str) + else json.dumps(exc.detail, default=_safe_str), + failure_category=failure_category_of(exc), file_data=request_dict.get("file_data", None), ) except UnstructuredIngestError as exc: logger.error( - f"UnstructuredIngestError: {str(exc)} (status_code={exc.status_code})", + f"UnstructuredIngestError: {_safe_str(exc)} " + f"(status_code={_error_attr(exc, 'status_code')})", exc_info=True, ) return InvokeResponse( usage=usage, message_channels=message_channels, 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), + status_code=status_code_of(exc), + status_code_text=_safe_str(exc), + failure_category=failure_category_of(exc), file_data=request_dict.get("file_data", None), ) except Exception as invoke_error: - logger.error(f"failed to invoke plugin: {invoke_error}", exc_info=True) + logger.error(f"failed to invoke plugin: {_safe_str(invoke_error)}", exc_info=True) return InvokeResponse( usage=usage, message_channels=message_channels, filedata_meta=filedata_meta_model.model_validate(filedata_meta.model_dump()), - 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}", + status_code=status_code_of(invoke_error), + status_code_text=f"[{type(invoke_error).__name__}] {_safe_str(invoke_error)}", + failure_category=failure_category_of(invoke_error), file_data=request_dict.get("file_data", None), ) @@ -287,9 +346,7 @@ async def run_job_with_body(request: BaseModel) -> ResponseType: @fastapi_app.post("/invoke", response_model=InvokeResponse) async def run_job(request: Optional[input_schema_model] = None) -> ResponseType: - return await run_job_with_body( - request if request is not None else input_schema_model() - ) + return await run_job_with_body(request if request is not None else input_schema_model()) elif input_schema_model.model_fields: @@ -318,6 +375,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 +390,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: