From 915a1cabcdcea66e34e14d6ebb6cedf1bffcdaa0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 17 Aug 2026 15:44:06 -0700 Subject: [PATCH] feat(proxy): add Amazon Comprehend Medical passthrough provider --- litellm/proxy/_types.py | 1 + .../billable_request_metrics_middleware.py | 1 + .../llm_passthrough_endpoints.py | 125 ++++++++++++ ...end_medical_passthrough_logging_handler.py | 102 ++++++++++ .../pass_through_endpoints/success_handler.py | 28 +++ ...est_billable_request_metrics_middleware.py | 3 + ...end_medical_passthrough_logging_handler.py | 134 +++++++++++++ .../test_llm_pass_through_endpoints.py | 186 ++++++++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 100 ++++++++++ 9 files changed, 680 insertions(+) create mode 100644 litellm/proxy/pass_through_endpoints/llm_provider_handlers/comprehend_medical_passthrough_logging_handler.py create mode 100644 tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_comprehend_medical_passthrough_logging_handler.py diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 22cc961a7d2..9a118f57f39 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -451,6 +451,7 @@ class LiteLLMRoutes(enum.Enum): mapped_pass_through_routes = [ "/bedrock", + "/comprehendmedical", "/vertex-ai", "/vertex_ai", "/cohere", diff --git a/litellm/proxy/middleware/billable_request_metrics_middleware.py b/litellm/proxy/middleware/billable_request_metrics_middleware.py index 9824f33797c..ac119e81d9c 100644 --- a/litellm/proxy/middleware/billable_request_metrics_middleware.py +++ b/litellm/proxy/middleware/billable_request_metrics_middleware.py @@ -92,6 +92,7 @@ _LLM_ROUTE_EXACT: Final[tuple[str, ...]] = ( "/v1/messages", "/interactions", # Google Interactions create; /{id} reads and /cancel do not match "/v1beta/interactions", + "/comprehendmedical", # AWS-SDK-shaped passthrough: the operation rides in the X-Amz-Target header ) # Provider passthrough prefixes (e.g. /bedrock/..., /vertex-ai/...) carry real diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 8cdcdc07547..635767f4db7 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -9,6 +9,7 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc. import json import os import re +from types import MappingProxyType from typing import Annotated, Any, Final, cast import httpx @@ -1079,6 +1080,130 @@ async def bedrock_proxy_route( return received_value +COMPREHEND_MEDICAL_TARGET_PREFIX: Final = "ComprehendMedical_20181030" + + +def _resolve_comprehend_medical_region() -> str | None: + region_candidates: Final = ( + get_secret_str(secret_name="AWS_REGION_NAME"), + get_secret_str(secret_name="AWS_REGION"), + get_secret_str(secret_name="AWS_DEFAULT_REGION"), + ) + return next((region for region in region_candidates if region), None) + + +@router.post( + "/comprehendmedical/{operation}", + tags=["AWS Comprehend Medical Pass-through", "pass-through"], # mutable-ok: fastapi route tags must be a list +) +async def comprehend_medical_proxy_route( + operation: str, + request: Request, + fastapi_response: Response, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +): + """ + Pass-through for Amazon Comprehend Medical, e.g. `POST /comprehendmedical/DetectEntitiesV2`. + + The request body is forwarded as-is to the AWS JSON 1.1 API and signed with SigV4 + using the proxy's AWS credentials. + + [Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical) + """ + try: + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest + from botocore.credentials import Credentials + except ImportError: + raise ImportError("Missing boto3 to call comprehendmedical. Run 'pip install boto3'.") + + from .llm_provider_handlers.comprehend_medical_passthrough_logging_handler import ( + COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS, + ) + + if operation not in COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS: + raise HTTPException( + status_code=400, + detail=( + f"Unsupported Comprehend Medical operation: {operation}. " + f"Supported operations: {', '.join(sorted(COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS))}" + ), + ) + + aws_region_name: Final = _resolve_comprehend_medical_region() + if aws_region_name is None: + raise HTTPException( + status_code=400, + detail="AWS region not found. Set AWS_REGION_NAME in the proxy environment.", + ) + + try: + data: Final = await request.json() + except Exception as e: + raise HTTPException(status_code=400, detail=str(e)) + + if not isinstance(data, dict): + raise HTTPException(status_code=400, detail="Request body must be a JSON object") + if "stream" in data: + raise HTTPException(status_code=400, detail="'stream' is not a Comprehend Medical request member") + + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + credentials: Final[Credentials] = BaseAWSLLM().get_credentials(aws_region_name=aws_region_name) + sigv4: Final = SigV4Auth(credentials, "comprehendmedical", aws_region_name) + headers: Final = MappingProxyType( + { + "Content-Type": "application/x-amz-json-1.1", + "X-Amz-Target": f"{COMPREHEND_MEDICAL_TARGET_PREFIX}.{operation}", + } + ) + target_url: Final = f"https://comprehendmedical.{aws_region_name}.amazonaws.com/" + _request: Final = AWSRequest(method="POST", url=target_url, data=json.dumps(data), headers=headers) + sigv4.add_auth(_request) + prepped: Final = _request.prepare() + + endpoint_func: Final = create_pass_through_route( + endpoint=operation, + target=str(prepped.url), + custom_headers=prepped.headers, + custom_llm_provider="comprehendmedical", + ) + setattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, data) + setattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, prepped.body) + return await endpoint_func(request, fastapi_response, user_api_key_dict) + + +@router.post( + "/comprehendmedical", + tags=["AWS Comprehend Medical Pass-through", "pass-through"], # mutable-ok: fastapi route tags must be a list +) +async def comprehend_medical_sdk_proxy_route( + request: Request, + fastapi_response: Response, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +): + """ + AWS-SDK-shaped pass-through for Amazon Comprehend Medical: point the SDK's + `endpoint_url` at `/comprehendmedical` and the operation is read from the + `X-Amz-Target` header, per the AWS JSON 1.1 protocol. + + [Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical) + """ + target_header: Final = request.headers.get("x-amz-target", "") + target_prefix, _, operation = target_header.partition(".") + if target_prefix != COMPREHEND_MEDICAL_TARGET_PREFIX or not operation: + raise HTTPException( + status_code=400, + detail=f"Expected an X-Amz-Target header of the form {COMPREHEND_MEDICAL_TARGET_PREFIX}.", + ) + return await comprehend_medical_proxy_route( + operation=operation, + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + ) + + def _resolve_vertex_model_from_router( model_id: str, llm_router: litellm.Router | None, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/comprehend_medical_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/comprehend_medical_passthrough_logging_handler.py new file mode 100644 index 00000000000..0d82cabdf36 --- /dev/null +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/comprehend_medical_passthrough_logging_handler.py @@ -0,0 +1,102 @@ +import math +from collections.abc import Mapping +from datetime import datetime +from types import MappingProxyType +from typing import Final + +import httpx + +from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.litellm_logging import ( + get_standard_logging_object_payload, +) +from litellm.proxy._types import PassThroughEndpointLoggingTypedDict +from litellm.types.utils import StandardPassThroughResponseObject + +COMPREHEND_MEDICAL_CHARS_PER_UNIT: Final = 100 +COMPREHEND_MEDICAL_COST_PER_UNIT_USD: Final[Mapping[str, float]] = MappingProxyType( + { + "DetectEntitiesV2": 0.01, + "DetectPHI": 0.0014, + "InferICD10CM": 0.0005, + "InferRxNorm": 0.00025, + "InferSNOMEDCT": 0.0075, + } +) +COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS: Final = frozenset(COMPREHEND_MEDICAL_COST_PER_UNIT_USD) + + +class ComprehendMedicalPassthroughLoggingHandler: + @staticmethod + def _operation_from_response(httpx_response: httpx.Response) -> str: + target: Final = httpx_response.request.headers.get("x-amz-target", "") + return target.split(".")[-1] + + @staticmethod + def get_cost_for_operation(operation: str, text: str) -> float: + cost_per_unit: Final = COMPREHEND_MEDICAL_COST_PER_UNIT_USD.get(operation) + if cost_per_unit is None: + return 0.0 + units: Final = max(1, math.ceil(len(text) / COMPREHEND_MEDICAL_CHARS_PER_UNIT)) + return units * cost_per_unit + + @staticmethod + def comprehend_medical_passthrough_handler( + httpx_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + url_route: str, + result: str, + start_time: datetime, + end_time: datetime, + cache_hit: bool, + request_body: Mapping[str, object], + **kwargs: object, # kwargs-ok: the passthrough logging dispatch forwards shared logging kwargs to every handler + ) -> PassThroughEndpointLoggingTypedDict: + """ + Prices a Comprehend Medical sync operation from the request text length + (billed per started 100-character unit, 1-unit minimum) and records + model, provider, and cost on the logging payload. + """ + try: + operation: Final = ComprehendMedicalPassthroughLoggingHandler._operation_from_response(httpx_response) + text: Final = request_body.get("Text") + response_cost: Final = ComprehendMedicalPassthroughLoggingHandler.get_cost_for_operation( + operation=operation, + text=text if isinstance(text, str) else "", + ) + model_name: Final = f"comprehendmedical/{operation}" + + updated_kwargs: Final = { # mutable-ok: the logging pipeline requires a plain kwargs dict + **kwargs, + "model": model_name, + "custom_llm_provider": "comprehendmedical", + "response_cost": response_cost, + } + logging_obj.model_call_details.update( + model=model_name, + custom_llm_provider="comprehendmedical", + response_cost=response_cost, + ) + + standard_logging_object: Final = get_standard_logging_object_payload( + kwargs=updated_kwargs, + init_response_obj=StandardPassThroughResponseObject(response=result), + start_time=start_time, + end_time=end_time, + logging_obj=logging_obj, + status="success", + ) + + handler_payload: Final[PassThroughEndpointLoggingTypedDict] = { + "result": StandardPassThroughResponseObject(response=result), + "kwargs": {**updated_kwargs, "standard_logging_object": standard_logging_object}, + } + except Exception as e: + verbose_proxy_logger.exception("Error in Comprehend Medical passthrough logging handler: %s", e) + fallback_payload: Final[PassThroughEndpointLoggingTypedDict] = { + "result": StandardPassThroughResponseObject(response=result), + "kwargs": kwargs, + } + return fallback_payload + return handler_payload diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 34286b203c7..749784c1bf5 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -236,6 +236,26 @@ class PassThroughEndpointLogging: ) standard_logging_response_object = cursor_passthrough_logging_handler_result["result"] kwargs = cursor_passthrough_logging_handler_result["kwargs"] + elif self.is_comprehend_medical_route(url_route, custom_llm_provider): + from .llm_provider_handlers.comprehend_medical_passthrough_logging_handler import ( + ComprehendMedicalPassthroughLoggingHandler, + ) + + comprehend_medical_handler_result: Final = ( + ComprehendMedicalPassthroughLoggingHandler.comprehend_medical_passthrough_handler( + httpx_response=httpx_response, + logging_obj=logging_obj, + url_route=url_route, + result=result, + start_time=start_time, + end_time=end_time, + cache_hit=cache_hit, + request_body=request_body, + **kwargs, + ) + ) + standard_logging_response_object = comprehend_medical_handler_result["result"] # rebind-ok: elif-chain + kwargs = comprehend_medical_handler_result["kwargs"] # rebind-ok: elif-chain contract elif self.is_vertex_ai_live_route(url_route): from .llm_provider_handlers.vertex_ai_live_passthrough_logging_handler import ( VertexAILivePassthroughLoggingHandler, @@ -364,6 +384,14 @@ class PassThroughEndpointLogging: return True return False + def is_comprehend_medical_route(self, url_route: str, custom_llm_provider: str | None = None) -> bool: + if custom_llm_provider == "comprehendmedical": + return True + hostname: Final = urlparse(url_route).hostname + if hostname is None: + return False + return hostname.startswith("comprehendmedical.") and hostname.endswith(".amazonaws.com") + def is_langfuse_route(self, url_route: str): parsed_url: Final = urlparse(url_route) for route in self.TRACKED_LANGFUSE_ROUTES: diff --git a/tests/test_litellm/proxy/middleware/test_billable_request_metrics_middleware.py b/tests/test_litellm/proxy/middleware/test_billable_request_metrics_middleware.py index ff7a24db832..9c61412bd6e 100644 --- a/tests/test_litellm/proxy/middleware/test_billable_request_metrics_middleware.py +++ b/tests/test_litellm/proxy/middleware/test_billable_request_metrics_middleware.py @@ -113,6 +113,9 @@ def test_is_pure_asgi_not_base_http_middleware(): ("/cohere/v2/chat", (BillableCategory.LLM, "/cohere")), # Passthrough inference bills under its provider prefix ("/anthropic/v1/messages", (BillableCategory.LLM, "/anthropic")), + # Bare AWS-SDK-shaped route carries the operation in X-Amz-Target and writes SpendLogs + ("/comprehendmedical", (BillableCategory.LLM, "/comprehendmedical")), + ("/comprehendmedical/DetectEntitiesV2", (BillableCategory.LLM, "/comprehendmedical")), ("/mcp", (BillableCategory.MCP, "/mcp")), ("/mcp/", (BillableCategory.MCP, "/mcp")), ("/mcp/tools/list", (BillableCategory.MCP, "/mcp")), diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_comprehend_medical_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_comprehend_medical_passthrough_logging_handler.py new file mode 100644 index 00000000000..0bc45c046bd --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_comprehend_medical_passthrough_logging_handler.py @@ -0,0 +1,134 @@ +import os +import sys +from datetime import datetime +from unittest.mock import MagicMock + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.proxy.pass_through_endpoints.llm_provider_handlers.comprehend_medical_passthrough_logging_handler import ( + ComprehendMedicalPassthroughLoggingHandler, +) +from litellm.proxy.pass_through_endpoints.success_handler import ( + PassThroughEndpointLogging, +) + + +def _make_response(operation: str) -> httpx.Response: + request = httpx.Request( + "POST", + "https://comprehendmedical.us-east-1.amazonaws.com/", + headers={"X-Amz-Target": f"ComprehendMedical_20181030.{operation}"}, + ) + return httpx.Response(200, request=request, text='{"Entities": []}') + + +def _make_logging_obj() -> MagicMock: + logging_obj = MagicMock() + logging_obj.litellm_call_id = "test-call-id" + logging_obj.model_call_details = {} + return logging_obj + + +class TestComprehendMedicalCost: + @pytest.mark.parametrize( + "operation,text,expected", + [ + ("DetectEntitiesV2", "x" * 250, 0.03), + ("DetectEntitiesV2", "x" * 100, 0.01), + ("DetectPHI", "", 0.0014), + ("DetectPHI", "x" * 101, 0.0028), + ("InferICD10CM", "x" * 100, 0.0005), + ("InferRxNorm", "x" * 150, 0.0005), + ("InferSNOMEDCT", "x", 0.0075), + ("StartEntitiesDetectionV2Job", "x" * 1000, 0.0), + ], + ) + def test_cost_per_started_100_char_unit(self, operation, text, expected): + assert ComprehendMedicalPassthroughLoggingHandler.get_cost_for_operation( + operation=operation, text=text + ) == pytest.approx(expected) + + +class TestComprehendMedicalPassthroughHandler: + def test_records_model_provider_and_cost(self): + logging_obj = _make_logging_obj() + + handler_result = ComprehendMedicalPassthroughLoggingHandler.comprehend_medical_passthrough_handler( + httpx_response=_make_response("DetectEntitiesV2"), + logging_obj=logging_obj, + url_route="https://comprehendmedical.us-east-1.amazonaws.com/", + result='{"Entities": []}', + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={"Text": "x" * 250}, + ) + + assert handler_result["result"] == {"response": '{"Entities": []}'} + assert handler_result["kwargs"]["model"] == "comprehendmedical/DetectEntitiesV2" + assert handler_result["kwargs"]["custom_llm_provider"] == "comprehendmedical" + assert handler_result["kwargs"]["response_cost"] == pytest.approx(0.03) + assert "standard_logging_object" in handler_result["kwargs"] + assert logging_obj.model_call_details["model"] == "comprehendmedical/DetectEntitiesV2" + assert logging_obj.model_call_details["custom_llm_provider"] == "comprehendmedical" + assert logging_obj.model_call_details["response_cost"] == pytest.approx(0.03) + + def test_missing_text_bills_one_unit_minimum(self): + logging_obj = _make_logging_obj() + + handler_result = ComprehendMedicalPassthroughLoggingHandler.comprehend_medical_passthrough_handler( + httpx_response=_make_response("DetectPHI"), + logging_obj=logging_obj, + url_route="https://comprehendmedical.us-east-1.amazonaws.com/", + result="{}", + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={}, + ) + + assert handler_result["kwargs"]["model"] == "comprehendmedical/DetectPHI" + assert handler_result["kwargs"]["response_cost"] == pytest.approx(0.0014) + + +class TestIsComprehendMedicalRoute: + def test_matches_by_hostname(self): + assert PassThroughEndpointLogging().is_comprehend_medical_route( + "https://comprehendmedical.us-east-1.amazonaws.com/", None + ) + + def test_matches_by_provider_tag(self): + assert PassThroughEndpointLogging().is_comprehend_medical_route("https://example.com/", "comprehendmedical") + + def test_does_not_match_other_aws_hosts(self): + assert not PassThroughEndpointLogging().is_comprehend_medical_route( + "https://bedrock-runtime.us-east-1.amazonaws.com/model/x/converse", None + ) + + def test_does_not_match_lookalike_hosts_outside_aws(self): + assert not PassThroughEndpointLogging().is_comprehend_medical_route("https://comprehendmedical.evil.com/", None) + + +class TestNormalizeDispatch: + def test_normalize_routes_to_comprehend_medical_handler(self): + logging_obj = _make_logging_obj() + + normalized = PassThroughEndpointLogging().normalize_llm_passthrough_logging_payload( + httpx_response=_make_response("DetectPHI"), + response_body={"Entities": []}, + request_body={"Text": "John Smith"}, + logging_obj=logging_obj, + url_route="https://comprehendmedical.us-east-1.amazonaws.com/", + result='{"Entities": []}', + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + custom_llm_provider="comprehendmedical", + ) + + assert normalized["standard_logging_response_object"] == {"response": '{"Entities": []}'} + assert normalized["kwargs"]["model"] == "comprehendmedical/DetectPHI" + assert normalized["kwargs"]["response_cost"] == pytest.approx(0.0014) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 9e6a2d42757..050070e2fcf 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -3471,3 +3471,189 @@ class TestAzureProxyRouteServiceLevelIndexCreate: ) mock_handler.assert_awaited_once() + + +class TestComprehendMedicalProxyRoute: + def _mock_request(self, body: object) -> Mock: + mock_request = Mock() + mock_request.method = "POST" + mock_request.json = AsyncMock(return_value=body) + return mock_request + + @pytest.mark.asyncio + async def test_signs_and_forwards_detect_entities_v2(self): + from botocore.credentials import Credentials + + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + comprehend_medical_proxy_route, + ) + from litellm.types.passthrough_endpoints.pass_through_endpoints import ( + LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, + LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, + ) + + request_body = {"Text": "Patient was prescribed 40mg atorvastatin daily."} + mock_request = self._mock_request(request_body) + mock_endpoint_func = AsyncMock(return_value={"Entities": []}) + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str", + side_effect=lambda secret_name: "us-east-1" if secret_name == "AWS_REGION_NAME" else None, + ), + patch( + "litellm.llms.bedrock.base_aws_llm.BaseAWSLLM.get_credentials", + return_value=Credentials("test-access-key", "test-secret-key"), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route", + return_value=mock_endpoint_func, + ) as mock_create_route, + ): + result = await comprehend_medical_proxy_route( + operation="DetectEntitiesV2", + request=mock_request, + fastapi_response=Mock(), + user_api_key_dict=Mock(), + ) + + assert result == {"Entities": []} + call_kwargs = mock_create_route.call_args.kwargs + assert call_kwargs["target"] == "https://comprehendmedical.us-east-1.amazonaws.com/" + assert call_kwargs["custom_llm_provider"] == "comprehendmedical" + assert "_forward_headers" not in call_kwargs + signed_headers = dict(call_kwargs["custom_headers"]) + assert signed_headers["X-Amz-Target"] == "ComprehendMedical_20181030.DetectEntitiesV2" + assert signed_headers["Content-Type"] == "application/x-amz-json-1.1" + assert signed_headers["Authorization"].startswith("AWS4-HMAC-SHA256") + assert "/comprehendmedical/aws4_request" in signed_headers["Authorization"] + assert getattr(mock_request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY) == request_body + assert json.loads(getattr(mock_request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY)) == request_body + mock_endpoint_func.assert_awaited_once() + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "operation", + [ + "Detect-Entities", + "Detect/../secrets", + "", + "a" * 200, + "DetectEntities", + "StartEntitiesDetectionV2Job", + ], + ) + async def test_rejects_unsupported_operations(self, operation): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + comprehend_medical_proxy_route, + ) + + with pytest.raises(HTTPException) as exc_info: + await comprehend_medical_proxy_route( + operation=operation, + request=self._mock_request({"Text": "hi"}), + fastapi_response=Mock(), + user_api_key_dict=Mock(), + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + @pytest.mark.parametrize("body", [{"Text": "hi", "stream": True}, {"Text": "hi", "stream": False}, ["Text"]]) + async def test_rejects_stream_key_and_non_object_bodies(self, body): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + comprehend_medical_proxy_route, + ) + + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str", + return_value="us-east-1", + ): + with pytest.raises(HTTPException) as exc_info: + await comprehend_medical_proxy_route( + operation="DetectEntitiesV2", + request=self._mock_request(body), + fastapi_response=Mock(), + user_api_key_dict=Mock(), + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_missing_region_returns_400(self): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + comprehend_medical_proxy_route, + ) + + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str", + return_value=None, + ): + with pytest.raises(HTTPException) as exc_info: + await comprehend_medical_proxy_route( + operation="DetectPHI", + request=self._mock_request({"Text": "hi"}), + fastapi_response=Mock(), + user_api_key_dict=Mock(), + ) + assert exc_info.value.status_code == 400 + + def test_comprehendmedical_is_a_mapped_pass_through_route(self): + from litellm.proxy._types import LiteLLMRoutes + + assert "/comprehendmedical" in LiteLLMRoutes.mapped_pass_through_routes.value + + @pytest.mark.asyncio + async def test_sdk_route_reads_operation_from_x_amz_target(self): + from botocore.credentials import Credentials + + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + comprehend_medical_sdk_proxy_route, + ) + + mock_request = self._mock_request({"Text": "hi"}) + mock_request.headers = {"x-amz-target": "ComprehendMedical_20181030.DetectPHI"} + mock_endpoint_func = AsyncMock(return_value={"Entities": []}) + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str", + side_effect=lambda secret_name: "us-east-1" if secret_name == "AWS_REGION_NAME" else None, + ), + patch( + "litellm.llms.bedrock.base_aws_llm.BaseAWSLLM.get_credentials", + return_value=Credentials("test-access-key", "test-secret-key"), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route", + return_value=mock_endpoint_func, + ) as mock_create_route, + ): + result = await comprehend_medical_sdk_proxy_route( + request=mock_request, + fastapi_response=Mock(), + user_api_key_dict=Mock(), + ) + + assert result == {"Entities": []} + signed_headers = dict(mock_create_route.call_args.kwargs["custom_headers"]) + assert signed_headers["X-Amz-Target"] == "ComprehendMedical_20181030.DetectPHI" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "target_header", + ["", "ComprehendMedical_20181030", "WrongService.DetectPHI", "ComprehendMedical_20181030."], + ) + async def test_sdk_route_rejects_bad_x_amz_target(self, target_header): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + comprehend_medical_sdk_proxy_route, + ) + + mock_request = self._mock_request({"Text": "hi"}) + mock_request.headers = {"x-amz-target": target_header} + + with pytest.raises(HTTPException) as exc_info: + await comprehend_medical_sdk_proxy_route( + request=mock_request, + fastapi_response=Mock(), + user_api_key_dict=Mock(), + ) + assert exc_info.value.status_code == 400 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 2bc21735248..996bd8d513a 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -2064,6 +2064,55 @@ export interface paths { patch?: never; trace?: never; }; + "/comprehendmedical": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** + * Comprehend Medical Sdk Proxy Route + * @description AWS-SDK-shaped pass-through for Amazon Comprehend Medical: point the SDK's + * `endpoint_url` at `/comprehendmedical` and the operation is read from the + * `X-Amz-Target` header, per the AWS JSON 1.1 protocol. + * + * [Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical) + */ + post: operations["comprehend_medical_sdk_proxy_route_comprehendmedical_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/comprehendmedical/{operation}": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** + * Comprehend Medical Proxy Route + * @description Pass-through for Amazon Comprehend Medical, e.g. `POST /comprehendmedical/DetectEntitiesV2`. + * + * The request body is forwarded as-is to the AWS JSON 1.1 API and signed with SigV4 + * using the proxy's AWS credentials. + * + * [Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical) + */ + post: operations["comprehend_medical_proxy_route_comprehendmedical__operation__post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/config/callback/delete": { parameters: { query?: never; @@ -39444,6 +39493,57 @@ export interface operations { }; }; }; + comprehend_medical_sdk_proxy_route_comprehendmedical_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; + comprehend_medical_proxy_route_comprehendmedical__operation__post: { + parameters: { + query?: never; + header?: never; + path: { + operation: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; delete_callback_config_callback_delete_post: { parameters: { query?: never;