1import math
2from collections.abc import Mapping
3from datetime import datetime
4from types import MappingProxyType
5from typing import Final
6
7import httpx
8
9from litellm._logging import verbose_proxy_logger
10from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
11from litellm.litellm_core_utils.litellm_logging import (
12 get_standard_logging_object_payload,
13)
14from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
15from litellm.types.utils import StandardPassThroughResponseObject
16
17COMPREHEND_MEDICAL_CHARS_PER_UNIT: Final = 100
18COMPREHEND_MEDICAL_COST_PER_UNIT_USD: Final[Mapping[str, float]] = MappingProxyType(
19 {
20 "DetectEntitiesV2": 0.01,
21 "DetectPHI": 0.0014,
22 "InferICD10CM": 0.0005,
23 "InferRxNorm": 0.00025,
24 "InferSNOMEDCT": 0.0075,
25 }
26)
27COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS: Final = frozenset(COMPREHEND_MEDICAL_COST_PER_UNIT_USD)
28
29
30class ComprehendMedicalPassthroughLoggingHandler:
31 @staticmethod
32 def _operation_from_response(httpx_response: httpx.Response) -> str:
33 target: Final = httpx_response.request.headers.get("x-amz-target", "")
34 return target.split(".")[-1]
35
36 @staticmethod
37 def get_cost_for_operation(operation: str, text: str) -> float:
38 cost_per_unit: Final = COMPREHEND_MEDICAL_COST_PER_UNIT_USD.get(operation)
39 if cost_per_unit is None:
40 return 0.0
41 units: Final = max(1, math.ceil(len(text) / COMPREHEND_MEDICAL_CHARS_PER_UNIT))
42 return units * cost_per_unit
43
44 @staticmethod
45 def comprehend_medical_passthrough_handler(
46 httpx_response: httpx.Response,
47 logging_obj: LiteLLMLoggingObj,
48 url_route: str,
49 result: str,
50 start_time: datetime,
51 end_time: datetime,
52 cache_hit: bool,
53 request_body: Mapping[str, object],
54 **kwargs: object, # kwargs-ok: the passthrough logging dispatch forwards shared logging kwargs to every handler
55 ) -> PassThroughEndpointLoggingTypedDict:
56 """
57 Prices a Comprehend Medical sync operation from the request text length
58 (billed per started 100-character unit, 1-unit minimum) and records
59 model, provider, and cost on the logging payload.
60 """
61 try:
62 operation: Final = ComprehendMedicalPassthroughLoggingHandler._operation_from_response(httpx_response)
63 text: Final = request_body.get("Text")
64 response_cost: Final = ComprehendMedicalPassthroughLoggingHandler.get_cost_for_operation(
65 operation=operation,
66 text=text if isinstance(text, str) else "",
67 )
68 model_name: Final = f"comprehendmedical/{operation}"
69
70 updated_kwargs: Final = { # mutable-ok: the logging pipeline requires a plain kwargs dict
71 **kwargs,
72 "model": model_name,
73 "custom_llm_provider": "comprehendmedical",
74 "response_cost": response_cost,
75 }
76 logging_obj.model_call_details.update(
77 model=model_name,
78 custom_llm_provider="comprehendmedical",
79 response_cost=response_cost,
80 )
81
82 standard_logging_object: Final = get_standard_logging_object_payload(
83 kwargs=updated_kwargs,
84 init_response_obj=StandardPassThroughResponseObject(response=result),
85 start_time=start_time,
86 end_time=end_time,
87 logging_obj=logging_obj,
88 status="success",
89 )
90
91 handler_payload: Final[PassThroughEndpointLoggingTypedDict] = {
92 "result": StandardPassThroughResponseObject(response=result),
93 "kwargs": {**updated_kwargs, "standard_logging_object": standard_logging_object},
94 }
95 except Exception as e:
96 verbose_proxy_logger.exception("Error in Comprehend Medical passthrough logging handler: %s", e)
97 fallback_payload: Final[PassThroughEndpointLoggingTypedDict] = {
98 "result": StandardPassThroughResponseObject(response=result),
99 "kwargs": kwargs,
100 }
101 return fallback_payload
102 return handler_payload