From 7eecba6a8534559fe09cb947565beb11261426c6 Mon Sep 17 00:00:00 2001 From: Tim Elfrink Date: Wed, 17 Sep 2025 07:30:31 +0200 Subject: [PATCH] Implement AWS Bedrock CountTokens API support - Add support for both Converse and InvokeModel input formats - Implement endpoint handling in pass_through_endpoints - Add transformation logic for AWS Bedrock CountTokens API - Simplify model resolution using existing router patterns - Support token counting for messages and raw text inputs --- litellm/llms/bedrock/count_tokens/handler.py | 116 +++++++++ .../bedrock/count_tokens/transformation.py | 223 ++++++++++++++++++ .../llm_passthrough_endpoints.py | 83 +++++++ 3 files changed, 422 insertions(+) create mode 100644 litellm/llms/bedrock/count_tokens/handler.py create mode 100644 litellm/llms/bedrock/count_tokens/transformation.py diff --git a/litellm/llms/bedrock/count_tokens/handler.py b/litellm/llms/bedrock/count_tokens/handler.py new file mode 100644 index 00000000000..7904cb08d24 --- /dev/null +++ b/litellm/llms/bedrock/count_tokens/handler.py @@ -0,0 +1,116 @@ +""" +AWS Bedrock CountTokens API handler. + +Simplified handler leveraging existing LiteLLM Bedrock infrastructure. +""" + +from typing import Any, Dict + +from fastapi import HTTPException + +from litellm._logging import verbose_logger +from litellm.llms.bedrock.count_tokens.transformation import BedrockCountTokensConfig + + +class BedrockCountTokensHandler(BedrockCountTokensConfig): + """ + Simplified handler for AWS Bedrock CountTokens API requests. + + Uses existing LiteLLM infrastructure for authentication and request handling. + """ + + async def handle_count_tokens_request( + self, + request_data: Dict[str, Any], + litellm_params: Dict[str, Any], + resolved_model: str, + ) -> Dict[str, Any]: + """ + Handle a CountTokens request using existing LiteLLM patterns. + + Args: + request_data: The incoming request payload + litellm_params: LiteLLM configuration parameters + resolved_model: The actual model ID resolved from router + + Returns: + Dictionary containing token count response + """ + try: + # Validate the request + self.validate_count_tokens_request(request_data) + + verbose_logger.debug(f"Processing CountTokens request for resolved model: {resolved_model}") + + # Get AWS region using existing LiteLLM function + aws_region_name = self._get_aws_region_name( + optional_params=litellm_params, + model=resolved_model, + model_id=None, + ) + + verbose_logger.debug(f"Retrieved AWS region: {aws_region_name}") + + # Transform request to Bedrock format (supports both Converse and InvokeModel) + bedrock_request = self.transform_anthropic_to_bedrock_count_tokens( + request_data=request_data + ) + + verbose_logger.debug(f"Transformed request: {bedrock_request}") + + # Get endpoint URL using simplified function + endpoint_url = self.get_bedrock_count_tokens_endpoint(resolved_model, aws_region_name) + + verbose_logger.debug(f"Making request to: {endpoint_url}") + + # Use existing _sign_request method from BaseAWSLLM + headers = {"Content-Type": "application/json"} + signed_headers, signed_body = self._sign_request( + service_name="bedrock", + headers=headers, + optional_params=litellm_params, + request_data=bedrock_request, + api_base=endpoint_url, + model=resolved_model, + ) + + # Make HTTP request + import httpx + async with httpx.AsyncClient() as client: + response = await client.post( + endpoint_url, + headers=signed_headers, + content=signed_body, + timeout=30.0, + ) + + verbose_logger.debug(f"Response status: {response.status_code}") + + if response.status_code != 200: + error_text = response.text + verbose_logger.error(f"AWS Bedrock error: {error_text}") + raise HTTPException( + status_code=400, + detail={"error": f"AWS Bedrock error: {error_text}"} + ) + + bedrock_response = response.json() + + verbose_logger.debug(f"Bedrock response: {bedrock_response}") + + # Transform response back to expected format + final_response = self.transform_bedrock_response_to_anthropic(bedrock_response) + + verbose_logger.debug(f"Final response: {final_response}") + + return final_response + + except HTTPException: + # Re-raise HTTP exceptions as-is + raise + except Exception as e: + verbose_logger.error(f"Error in CountTokens handler: {str(e)}") + raise HTTPException( + status_code=500, + detail={"error": f"CountTokens processing error: {str(e)}"} + ) \ No newline at end of file diff --git a/litellm/llms/bedrock/count_tokens/transformation.py b/litellm/llms/bedrock/count_tokens/transformation.py new file mode 100644 index 00000000000..285c3baac00 --- /dev/null +++ b/litellm/llms/bedrock/count_tokens/transformation.py @@ -0,0 +1,223 @@ +""" +AWS Bedrock CountTokens API transformation logic. + +This module handles the transformation of requests from Anthropic Messages API format +to AWS Bedrock's CountTokens API format and vice versa. +""" + +from typing import Any, Dict, List + +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.bedrock.common_utils import BedrockModelInfo + + +class BedrockCountTokensConfig(BaseAWSLLM): + """ + Configuration and transformation logic for AWS Bedrock CountTokens API. + + AWS Bedrock CountTokens API Specification: + - Endpoint: POST /model/{modelId}/count-tokens + - Input formats: 'invokeModel' or 'converse' + - Response: {"inputTokens": } + """ + + def _detect_input_type(self, request_data: Dict[str, Any]) -> str: + """ + Detect whether to use 'converse' or 'invokeModel' input format. + + Args: + request_data: The original request data + + Returns: + 'converse' or 'invokeModel' + """ + # If the request has messages in the expected Anthropic format, use converse + if "messages" in request_data and isinstance(request_data["messages"], list): + return "converse" + + # For raw text or other formats, use invokeModel + # This handles cases where the input is prompt-based or already in raw Bedrock format + return "invokeModel" + + def transform_anthropic_to_bedrock_count_tokens( + self, + request_data: Dict[str, Any], + ) -> Dict[str, Any]: + """ + Transform request to Bedrock CountTokens format. + Supports both Converse and InvokeModel input types. + + Input (Anthropic format): + { + "model": "claude-3-5-sonnet", + "messages": [{"role": "user", "content": "Hello!"}] + } + + Output (Bedrock CountTokens format for Converse): + { + "input": { + "converse": { + "messages": [...], + "system": [...] (if present) + } + } + } + + Output (Bedrock CountTokens format for InvokeModel): + { + "input": { + "invokeModel": { + "body": "{...raw model input...}" + } + } + } + """ + input_type = self._detect_input_type(request_data) + + if input_type == "converse": + return self._transform_to_converse_format(request_data.get("messages", [])) + else: + return self._transform_to_invoke_model_format(request_data) + + def _transform_to_converse_format(self, messages: List[Dict[str, Any]]) -> Dict[str, Any]: + """Transform to Converse input format.""" + # Extract system messages if present + system_messages = [] + user_messages = [] + + for message in messages: + if message.get("role") == "system": + system_messages.append({"text": message.get("content", "")}) + else: + # Transform message content to Bedrock format + transformed_message = { + "role": message.get("role"), + "content": [] + } + + # Handle content - ensure it's in the correct array format + content = message.get("content", "") + if isinstance(content, str): + # String content -> convert to text block + transformed_message["content"].append({"text": content}) + elif isinstance(content, list): + # Already in blocks format - use as is + transformed_message["content"] = content + + user_messages.append(transformed_message) + + # Build the converse input format + converse_input = { + "messages": user_messages + } + + # Add system messages if present + if system_messages: + converse_input["system"] = system_messages + + # Build the complete request + return { + "input": { + "converse": converse_input + } + } + + def _transform_to_invoke_model_format(self, request_data: Dict[str, Any]) -> Dict[str, Any]: + """Transform to InvokeModel input format.""" + import json + + # For InvokeModel, we need to provide the raw body that would be sent to the model + # Remove the 'model' field from the body as it's not part of the model input + body_data = {k: v for k, v in request_data.items() if k != "model"} + + return { + "input": { + "invokeModel": { + "body": json.dumps(body_data) + } + } + } + + def get_bedrock_count_tokens_endpoint(self, model: str, aws_region_name: str) -> str: + """ + Construct the AWS Bedrock CountTokens API endpoint using existing LiteLLM functions. + + Args: + model: The resolved model ID from router lookup + aws_region_name: AWS region (e.g., "eu-west-1") + + Returns: + Complete endpoint URL for CountTokens API + """ + # Use existing LiteLLM function to get the base model ID (removes region prefix) + model_id = BedrockModelInfo.get_base_model(model) + + # Remove bedrock/ prefix if present + if model_id.startswith("bedrock/"): + model_id = model_id[8:] # Remove "bedrock/" prefix + + base_url = f"https://bedrock-runtime.{aws_region_name}.amazonaws.com" + endpoint = f"{base_url}/model/{model_id}/count-tokens" + + return endpoint + + + def transform_bedrock_response_to_anthropic(self, bedrock_response: Dict[str, Any]) -> Dict[str, Any]: + """ + Transform Bedrock CountTokens response to Anthropic format. + + Input (Bedrock response): + { + "inputTokens": 123 + } + + Output (Anthropic format): + { + "input_tokens": 123 + } + """ + input_tokens = bedrock_response.get("inputTokens", 0) + + return { + "input_tokens": input_tokens + } + + def validate_count_tokens_request(self, request_data: Dict[str, Any]) -> None: + """ + Validate the incoming count tokens request. + Supports both Converse and InvokeModel input formats. + + Args: + request_data: The request payload + + Raises: + ValueError: If the request is invalid + """ + if not request_data.get("model"): + raise ValueError("model parameter is required") + + input_type = self._detect_input_type(request_data) + + if input_type == "converse": + # Validate Converse format (messages-based) + messages = request_data.get("messages", []) + if not messages: + raise ValueError("messages parameter is required for Converse input") + + if not isinstance(messages, list): + raise ValueError("messages must be a list") + + for i, message in enumerate(messages): + if not isinstance(message, dict): + raise ValueError(f"Message {i} must be a dictionary") + + if "role" not in message: + raise ValueError(f"Message {i} must have a 'role' field") + + if "content" not in message: + raise ValueError(f"Message {i} must have a 'content' field") + else: + # For InvokeModel format, we need at least some content to count tokens + # The content structure varies by model, so we do minimal validation + if len(request_data) <= 1: # Only has 'model' field + raise ValueError("Request must contain content to count tokens") \ No newline at end of file diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 82c5b3e343d..a35877ee125 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -464,6 +464,78 @@ async def anthropic_proxy_route( return received_value +async def handle_bedrock_count_tokens( + endpoint: str, + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth, + request_body: Dict[str, Any], +) -> Dict[str, Any]: + """ + Handle AWS Bedrock CountTokens API requests. + + This function processes count_tokens endpoints like: + - /v1/messages/count_tokens + - /v1/messages/count-tokens + """ + from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler + from litellm.proxy.proxy_server import llm_router + + try: + # Initialize the handler + handler = BedrockCountTokensHandler() + + # Extract model from request body + model = request_body.get("model") + if not model: + raise HTTPException( + status_code=400, + detail={"error": "Model is required in request body"} + ) + + # Get model parameters from router + litellm_params = {"user_api_key_dict": user_api_key_dict} + resolved_model = model # Default fallback + + if llm_router: + deployments = llm_router.get_model_list(model_name=model) + if deployments and len(deployments) > 0: + # Get the first matching deployment + deployment = deployments[0] + model_litellm_params = deployment.get("litellm_params", {}) + + # Get the resolved model ID from the configuration + if "model" in model_litellm_params: + resolved_model = model_litellm_params["model"] + + # Copy all litellm_params - BaseAWSLLM will handle AWS credential discovery + for key, value in model_litellm_params.items(): + if key != "user_api_key_dict": # Don't overwrite user_api_key_dict + litellm_params[key] = value + + verbose_proxy_logger.debug(f"Count tokens litellm_params: {litellm_params}") + verbose_proxy_logger.debug(f"Resolved model: {resolved_model}") + + # Handle the count tokens request + result = await handler.handle_count_tokens_request( + request_data=request_body, + litellm_params=litellm_params, + resolved_model=resolved_model, + ) + + return result + + except HTTPException: + # Re-raise HTTP exceptions as-is + raise + except Exception as e: + verbose_proxy_logger.error(f"Error in handle_bedrock_count_tokens: {str(e)}") + raise HTTPException( + status_code=500, + detail={"error": f"CountTokens processing error: {str(e)}"} + ) + + async def bedrock_llm_proxy_route( endpoint: str, request: Request, @@ -489,6 +561,17 @@ async def bedrock_llm_proxy_route( ) request_body = await _read_request_body(request=request) + + # Special handling for count_tokens endpoints + if "count_tokens" in endpoint or "count-tokens" in endpoint: + return await handle_bedrock_count_tokens( + endpoint=endpoint, + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + request_body=request_body, + ) + data: Dict[str, Any] = {} base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) try: