diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index 9429b6dad43..f0b89615a0d 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -1999,203 +1999,13 @@ response = embedding( ### Advanced - [Pass model/provider-specific Params](https://docs.litellm.ai/docs/completion/provider_specific_params#proxy-usage) ## Image Generation -Use this for stable diffusion, and amazon nova canvas on bedrock + +See [Bedrock Image Generation](./bedrock_image_gen) for using Stable Diffusion and Amazon Nova Canvas models on Bedrock. -### Usage +## Rerank API - - - -```python -import os -from litellm import image_generation - -os.environ["AWS_ACCESS_KEY_ID"] = "" -os.environ["AWS_SECRET_ACCESS_KEY"] = "" -os.environ["AWS_REGION_NAME"] = "" - -response = image_generation( - prompt="A cute baby sea otter", - model="bedrock/stability.stable-diffusion-xl-v0", - ) -print(f"response: {response}") -``` - -**Set optional params** -```python -import os -from litellm import image_generation - -os.environ["AWS_ACCESS_KEY_ID"] = "" -os.environ["AWS_SECRET_ACCESS_KEY"] = "" -os.environ["AWS_REGION_NAME"] = "" - -response = image_generation( - prompt="A cute baby sea otter", - model="bedrock/stability.stable-diffusion-xl-v0", - ### OPENAI-COMPATIBLE ### - size="128x512", # width=128, height=512 - ### PROVIDER-SPECIFIC ### see `AmazonStabilityConfig` in bedrock.py for all params - seed=30 - ) -print(f"response: {response}") -``` - - - -1. Setup config.yaml - -```yaml -model_list: - - model_name: amazon.nova-canvas-v1:0 - litellm_params: - model: bedrock/amazon.nova-canvas-v1:0 - aws_region_name: "us-east-1" - aws_secret_access_key: my-key # OPTIONAL - all boto3 auth params supported - aws_secret_access_id: my-id # OPTIONAL - all boto3 auth params supported -``` - -2. Start proxy - -```bash -litellm --config /path/to/config.yaml -``` - -3. Test it! - -```bash -curl -L -X POST 'http://0.0.0.0:4000/v1/images/generations' \ --H 'Content-Type: application/json' \ --H 'Authorization: Bearer $LITELLM_VIRTUAL_KEY' \ --d '{ - "model": "amazon.nova-canvas-v1:0", - "prompt": "A cute baby sea otter" -}' -``` - - - - -### Using Inference Profiles with Image Generation - -For AWS Bedrock Application Inference Profiles with image generation, use the `model_id` parameter to specify the inference profile ARN: - - - - -```python -from litellm import image_generation - -response = image_generation( - model="bedrock/amazon.nova-canvas-v1:0", - model_id="arn:aws:bedrock:eu-west-1:000000000000:application-inference-profile/a0a0a0a0a0a0", - prompt="A cute baby sea otter" -) -print(f"response: {response}") -``` - - - - -```yaml -model_list: - - model_name: nova-canvas-inference-profile - litellm_params: - model: bedrock/amazon.nova-canvas-v1:0 - model_id: arn:aws:bedrock:eu-west-1:000000000000:application-inference-profile/a0a0a0a0a0a0 - aws_region_name: "eu-west-1" -``` - - - - -## Supported AWS Bedrock Image Generation Models - -| Model Name | Function Call | -|----------------------|---------------------------------------------| -| Stable Diffusion 3 - v0 | `embedding(model="bedrock/stability.stability.sd3-large-v1:0", prompt=prompt)` | -| Stable Diffusion - v0 | `embedding(model="bedrock/stability.stable-diffusion-xl-v0", prompt=prompt)` | -| Stable Diffusion - v0 | `embedding(model="bedrock/stability.stable-diffusion-xl-v1", prompt=prompt)` | - - -## Rerank API - -Use Bedrock's Rerank API in the Cohere `/rerank` format. - -Supported Cohere Rerank Params -- `model` - the foundation model ARN -- `query` - the query to rerank against -- `documents` - the list of documents to rerank -- `top_n` - the number of results to return - - - - -```python -from litellm import rerank -import os - -os.environ["AWS_ACCESS_KEY_ID"] = "" -os.environ["AWS_SECRET_ACCESS_KEY"] = "" -os.environ["AWS_REGION_NAME"] = "" - -response = rerank( - model="bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0", # provide the model ARN - get this here https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/bedrock/client/list_foundation_models.html - query="hello", - documents=["hello", "world"], - top_n=2, -) - -print(response) -``` - - - - -1. Setup config.yaml - -```yaml -model_list: - - model_name: bedrock-rerank - litellm_params: - model: bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0 - aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID - aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY - aws_region_name: os.environ/AWS_REGION_NAME -``` - -2. Start proxy server - -```bash -litellm --config config.yaml - -# RUNNING on http://0.0.0.0:4000 -``` - -3. Test it! - -```bash -curl http://0.0.0.0:4000/rerank \ - -H "Authorization: Bearer sk-1234" \ - -H "Content-Type: application/json" \ - -d '{ - "model": "bedrock-rerank", - "query": "What is the capital of the United States?", - "documents": [ - "Carson City is the capital city of the American state of Nevada.", - "The Commonwealth of the Northern Mariana Islands is a group of islands in the Pacific Ocean. Its capital is Saipan.", - "Washington, D.C. is the capital of the United States.", - "Capital punishment has existed in the United States since before it was a country." - ], - "top_n": 3 - - - }' -``` - - - +See [Bedrock Rerank](./bedrock_rerank) for using Bedrock's Rerank API in the Cohere `/rerank` format. ## Bedrock Application Inference Profile @@ -2490,38 +2300,6 @@ model_list: -Text to Image : -```bash -curl -L -X POST 'http://0.0.0.0:4000/v1/images/generations' \ --H 'Content-Type: application/json' \ --H 'Authorization: Bearer $LITELLM_VIRTUAL_KEY' \ --d '{ - "model": "amazon.nova-canvas-v1:0", - "prompt": "A cute baby sea otter" -}' -``` - -Color Guided Generation: -```bash -curl -L -X POST 'http://0.0.0.0:4000/v1/images/generations' \ --H 'Content-Type: application/json' \ --H 'Authorization: Bearer $LITELLM_VIRTUAL_KEY' \ --d '{ - "model": "amazon.nova-canvas-v1:0", - "prompt": "A cute baby sea otter", - "taskType": "COLOR_GUIDED_GENERATION", - "colorGuidedGenerationParams":{"colors":["#FFFFFF"]} -}' -``` - -| Model Name | Function Call | -|-------------------------|---------------------------------------------| -| Stable Diffusion 3 - v0 | `image_generation(model="bedrock/stability.stability.sd3-large-v1:0", prompt=prompt)` | -| Stable Diffusion - v0 | `image_generation(model="bedrock/stability.stable-diffusion-xl-v0", prompt=prompt)` | -| Stable Diffusion - v1 | `image_generation(model="bedrock/stability.stable-diffusion-xl-v1", prompt=prompt)` | -| Amazon Nova Canvas - v0 | `image_generation(model="bedrock/amazon.nova-canvas-v1:0", prompt=prompt)` | - - ### Passing an external BedrockRuntime.Client as a parameter - Completion() This is a deprecated flow. Boto3 is not async. And boto3.client does not let us make the http call through httpx. Pass in your aws params through the method above 👆. [See Auth Code](https://github.com/BerriAI/litellm/blob/55a20c7cce99a93d36a82bf3ae90ba3baf9a7f89/litellm/llms/bedrock_httpx.py#L284) [Add new auth flow](https://github.com/BerriAI/litellm/issues) diff --git a/docs/my-website/docs/providers/bedrock_batches.md b/docs/my-website/docs/providers/bedrock_batches.md index 57487f7d2c9..c262eef0e86 100644 --- a/docs/my-website/docs/providers/bedrock_batches.md +++ b/docs/my-website/docs/providers/bedrock_batches.md @@ -9,6 +9,7 @@ Use Amazon Bedrock Batch Inference API through LiteLLM. |----------|---------| | Description | Amazon Bedrock Batch Inference allows you to run inference on large datasets asynchronously | | Provider Doc | [AWS Bedrock Batch Inference ↗](https://docs.aws.amazon.com/bedrock/latest/userguide/batch-inference.html) | +| Cost Tracking | ✅ Supported | ## Overview diff --git a/docs/my-website/docs/providers/bedrock_embedding.md b/docs/my-website/docs/providers/bedrock_embedding.md index cd492084711..76c9606533e 100644 --- a/docs/my-website/docs/providers/bedrock_embedding.md +++ b/docs/my-website/docs/providers/bedrock_embedding.md @@ -2,11 +2,11 @@ ## Supported Embedding Models -| Provider | LiteLLM Route | AWS Documentation | -|----------|---------------|-------------------| -| Amazon Titan | `bedrock/amazon.*` | [Amazon Titan Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/titan-embedding-models.html) | -| Cohere | `bedrock/cohere.*` | [Cohere Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere-embed.html) | -| TwelveLabs | `bedrock/us.twelvelabs.*` | [TwelveLabs](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-twelvelabs.html) | +| Provider | LiteLLM Route | AWS Documentation | Cost Tracking | +|----------|---------------|-------------------|---------------| +| Amazon Titan | `bedrock/amazon.*` | [Amazon Titan Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/titan-embedding-models.html) | ✅ | +| Cohere | `bedrock/cohere.*` | [Cohere Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere-embed.html) | ✅ | +| TwelveLabs | `bedrock/us.twelvelabs.*` | [TwelveLabs](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-twelvelabs.html) | ✅ | ## Async Invoke Support diff --git a/docs/my-website/docs/providers/bedrock_image_gen.md b/docs/my-website/docs/providers/bedrock_image_gen.md new file mode 100644 index 00000000000..799c6d46437 --- /dev/null +++ b/docs/my-website/docs/providers/bedrock_image_gen.md @@ -0,0 +1,150 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# AWS Bedrock - Image Generation + +Use Bedrock for image generation with Stable Diffusion, Amazon Titan Image Generator, and Amazon Nova Canvas models. + +## Supported Models + +| Model Name | Function Call | Cost Tracking | +|-------------------------|---------------------------------------------|---------------| +| Stable Diffusion 3 - v0 | `image_generation(model="bedrock/stability.stability.sd3-large-v1:0", prompt=prompt)` | ✅ | +| Stable Diffusion - v0 | `image_generation(model="bedrock/stability.stable-diffusion-xl-v0", prompt=prompt)` | ✅ | +| Stable Diffusion - v1 | `image_generation(model="bedrock/stability.stable-diffusion-xl-v1", prompt=prompt)` | ✅ | +| Amazon Titan Image Generator - v1 | `image_generation(model="bedrock/amazon.titan-image-generator-v1", prompt=prompt)` | ✅ | +| Amazon Titan Image Generator - v2 | `image_generation(model="bedrock/amazon.titan-image-generator-v2:0", prompt=prompt)` | ✅ | +| Amazon Nova Canvas - v1 | `image_generation(model="bedrock/amazon.nova-canvas-v1:0", prompt=prompt)` | ✅ | + +## Usage + + + + +### Basic Usage + +```python +import os +from litellm import image_generation + +os.environ["AWS_ACCESS_KEY_ID"] = "" +os.environ["AWS_SECRET_ACCESS_KEY"] = "" +os.environ["AWS_REGION_NAME"] = "" + +response = image_generation( + prompt="A cute baby sea otter", + model="bedrock/stability.stable-diffusion-xl-v0", +) +print(f"response: {response}") +``` + +### Set Optional Parameters + +```python +import os +from litellm import image_generation + +os.environ["AWS_ACCESS_KEY_ID"] = "" +os.environ["AWS_SECRET_ACCESS_KEY"] = "" +os.environ["AWS_REGION_NAME"] = "" + +response = image_generation( + prompt="A cute baby sea otter", + model="bedrock/stability.stable-diffusion-xl-v0", + ### OPENAI-COMPATIBLE ### + size="128x512", # width=128, height=512 + ### PROVIDER-SPECIFIC ### see `AmazonStabilityConfig` in bedrock.py for all params + seed=30 +) +print(f"response: {response}") +``` + + + + +### 1. Setup config.yaml + +```yaml +model_list: + - model_name: amazon.nova-canvas-v1:0 + litellm_params: + model: bedrock/amazon.nova-canvas-v1:0 + aws_region_name: "us-east-1" + aws_secret_access_key: my-key # OPTIONAL - all boto3 auth params supported + aws_secret_access_id: my-id # OPTIONAL - all boto3 auth params supported +``` + +### 2. Start proxy + +```bash +litellm --config /path/to/config.yaml +``` + +### 3. Test it! + +**Text to Image:** + +```bash +curl -L -X POST 'http://0.0.0.0:4000/v1/images/generations' \ +-H 'Content-Type: application/json' \ +-H 'Authorization: Bearer $LITELLM_VIRTUAL_KEY' \ +-d '{ + "model": "amazon.nova-canvas-v1:0", + "prompt": "A cute baby sea otter" +}' +``` + +**Color Guided Generation:** + +```bash +curl -L -X POST 'http://0.0.0.0:4000/v1/images/generations' \ +-H 'Content-Type: application/json' \ +-H 'Authorization: Bearer $LITELLM_VIRTUAL_KEY' \ +-d '{ + "model": "amazon.nova-canvas-v1:0", + "prompt": "A cute baby sea otter", + "taskType": "COLOR_GUIDED_GENERATION", + "colorGuidedGenerationParams":{"colors":["#FFFFFF"]} +}' +``` + + + + +## Using Inference Profiles with Image Generation + +For AWS Bedrock Application Inference Profiles with image generation, use the `model_id` parameter to specify the inference profile ARN: + + + + +```python +from litellm import image_generation + +response = image_generation( + model="bedrock/amazon.nova-canvas-v1:0", + model_id="arn:aws:bedrock:eu-west-1:000000000000:application-inference-profile/a0a0a0a0a0a0", + prompt="A cute baby sea otter" +) +print(f"response: {response}") +``` + + + + +```yaml +model_list: + - model_name: nova-canvas-inference-profile + litellm_params: + model: bedrock/amazon.nova-canvas-v1:0 + model_id: arn:aws:bedrock:eu-west-1:000000000000:application-inference-profile/a0a0a0a0a0a0 + aws_region_name: "eu-west-1" +``` + + + + +## Authentication + +All standard Bedrock authentication methods are supported for image generation. See [Bedrock Authentication](./bedrock#boto3---authentication) for details. + diff --git a/docs/my-website/docs/providers/bedrock_rerank.md b/docs/my-website/docs/providers/bedrock_rerank.md new file mode 100644 index 00000000000..86745eb5125 --- /dev/null +++ b/docs/my-website/docs/providers/bedrock_rerank.md @@ -0,0 +1,94 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# AWS Bedrock - Rerank API + +Use Bedrock's Rerank API in the Cohere `/rerank` format. + +:::info Cost Tracking + +✅ **Cost tracking is supported** for Bedrock Rerank API calls. + +::: + +## Supported Parameters + +- `model` - the foundation model ARN +- `query` - the query to rerank against +- `documents` - the list of documents to rerank +- `top_n` - the number of results to return + +## Usage + + + + +```python +from litellm import rerank +import os + +os.environ["AWS_ACCESS_KEY_ID"] = "" +os.environ["AWS_SECRET_ACCESS_KEY"] = "" +os.environ["AWS_REGION_NAME"] = "" + +response = rerank( + model="bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0", # provide the model ARN - get this here https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/bedrock/client/list_foundation_models.html + query="hello", + documents=["hello", "world"], + top_n=2, +) + +print(response) +``` + + + + +### 1. Setup config.yaml + +```yaml +model_list: + - model_name: bedrock-rerank + litellm_params: + model: bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0 + aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID + aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY + aws_region_name: os.environ/AWS_REGION_NAME +``` + +### 2. Start proxy server + +```bash +litellm --config config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +### 3. Test it! + +```bash +curl http://0.0.0.0:4000/rerank \ + -H "Authorization: Bearer sk-1234" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "bedrock-rerank", + "query": "What is the capital of the United States?", + "documents": [ + "Carson City is the capital city of the American state of Nevada.", + "The Commonwealth of the Northern Mariana Islands is a group of islands in the Pacific Ocean. Its capital is Saipan.", + "Washington, D.C. is the capital of the United States.", + "Capital punishment has existed in the United States since before it was a country." + ], + "top_n": 3 + + + }' +``` + + + + +## Authentication + +All standard Bedrock authentication methods are supported for rerank. See [Bedrock Authentication](./bedrock#boto3---authentication) for details. + diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index f0e311c303f..b58c7e033fb 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -484,6 +484,8 @@ const sidebars = { items: [ "providers/bedrock", "providers/bedrock_embedding", + "providers/bedrock_image_gen", + "providers/bedrock_rerank", "providers/bedrock_agents", "providers/bedrock_batches", "providers/bedrock_vector_store", diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 32fde2d0987..6e8fc6d0cce 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -227,7 +227,9 @@ class OpenTelemetry(CustomLogger): PeriodicExportingMetricReader, ) - normalized_endpoint = self._normalize_otel_endpoint(self.config.endpoint, 'metrics') + normalized_endpoint = self._normalize_otel_endpoint( + self.config.endpoint, "metrics" + ) _metric_exporter = OTLPMetricExporter( endpoint=normalized_endpoint, headers=OpenTelemetry._get_headers_dictionary(self.config.headers), @@ -664,7 +666,9 @@ class OpenTelemetry(CustomLogger): # Get the resource from the logger provider logger_provider = get_logger_provider() - resource = getattr(logger_provider, '_resource', None) or _get_litellm_resource() + resource = ( + getattr(logger_provider, "_resource", None) or _get_litellm_resource() + ) parent_ctx = span.get_span_context() provider = (kwargs.get("litellm_params") or {}).get( @@ -1302,7 +1306,9 @@ class OpenTelemetry(CustomLogger): "OpenTelemetry: intiializing http exporter. Value of OTEL_EXPORTER: %s", self.OTEL_EXPORTER, ) - normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, 'traces') + normalized_endpoint = self._normalize_otel_endpoint( + self.OTEL_ENDPOINT, "traces" + ) return BatchSpanProcessor( OTLPSpanExporterHTTP( endpoint=normalized_endpoint, headers=_split_otel_headers @@ -1313,7 +1319,9 @@ class OpenTelemetry(CustomLogger): "OpenTelemetry: intiializing grpc exporter. Value of OTEL_EXPORTER: %s", self.OTEL_EXPORTER, ) - normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, 'traces') + normalized_endpoint = self._normalize_otel_endpoint( + self.OTEL_ENDPOINT, "traces" + ) return BatchSpanProcessor( OTLPSpanExporterGRPC( endpoint=normalized_endpoint, headers=_split_otel_headers @@ -1340,7 +1348,7 @@ class OpenTelemetry(CustomLogger): _split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS) # Normalize endpoint for logs - ensure it points to /v1/logs instead of /v1/traces - normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, 'logs') + normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "logs") verbose_logger.debug( "OpenTelemetry: Log endpoint normalized from %s to %s", @@ -1358,6 +1366,7 @@ class OpenTelemetry(CustomLogger): if self.OTEL_EXPORTER == "console": from opentelemetry.sdk._logs.export import ConsoleLogExporter + verbose_logger.debug( "OpenTelemetry: Using console log exporter. Value of OTEL_EXPORTER: %s", self.OTEL_EXPORTER, @@ -1368,7 +1377,10 @@ class OpenTelemetry(CustomLogger): or self.OTEL_EXPORTER == "http/protobuf" or self.OTEL_EXPORTER == "http/json" ): - from opentelemetry.exporter.otlp.proto.http._log_exporter import OTLPLogExporter + from opentelemetry.exporter.otlp.proto.http._log_exporter import ( + OTLPLogExporter, + ) + verbose_logger.debug( "OpenTelemetry: Using HTTP log exporter. Value of OTEL_EXPORTER: %s, endpoint: %s", self.OTEL_EXPORTER, @@ -1378,7 +1390,10 @@ class OpenTelemetry(CustomLogger): endpoint=normalized_endpoint, headers=_split_otel_headers ) elif self.OTEL_EXPORTER == "otlp_grpc" or self.OTEL_EXPORTER == "grpc": - from opentelemetry.exporter.otlp.proto.grpc._log_exporter import OTLPLogExporter + from opentelemetry.exporter.otlp.proto.grpc._log_exporter import ( + OTLPLogExporter, + ) + verbose_logger.debug( "OpenTelemetry: Using gRPC log exporter. Value of OTEL_EXPORTER: %s, endpoint: %s", self.OTEL_EXPORTER, @@ -1393,12 +1408,11 @@ class OpenTelemetry(CustomLogger): self.OTEL_EXPORTER, ) from opentelemetry.sdk._logs.export import ConsoleLogExporter + return ConsoleLogExporter() def _normalize_otel_endpoint( - self, - endpoint: Optional[str], - signal_type: str + self, endpoint: Optional[str], signal_type: str ) -> Optional[str]: """ Normalize the endpoint URL for a specific OpenTelemetry signal type. @@ -1431,37 +1445,37 @@ class OpenTelemetry(CustomLogger): return endpoint # Validate signal_type - valid_signals = {'traces', 'metrics', 'logs'} + valid_signals = {"traces", "metrics", "logs"} if signal_type not in valid_signals: verbose_logger.warning( "Invalid signal_type '%s' provided to _normalize_otel_endpoint. " "Valid values: %s. Returning endpoint unchanged.", signal_type, - valid_signals + valid_signals, ) return endpoint # Remove trailing slash - endpoint = endpoint.rstrip('/') + endpoint = endpoint.rstrip("/") # Check if endpoint already ends with the correct signal path - target_path = f'/v1/{signal_type}' + target_path = f"/v1/{signal_type}" if endpoint.endswith(target_path): return endpoint # Replace existing signal path with the target signal path other_signals = valid_signals - {signal_type} for other_signal in other_signals: - other_path = f'/v1/{other_signal}' + other_path = f"/v1/{other_signal}" if endpoint.endswith(other_path): - endpoint = endpoint.rsplit('/', 1)[0] + f'/{signal_type}' + endpoint = endpoint.rsplit("/", 1)[0] + f"/{signal_type}" return endpoint # No existing signal path found, append the target path - if not endpoint.endswith('/v1'): + if not endpoint.endswith("/v1"): endpoint = endpoint + target_path else: - endpoint = endpoint + f'/{signal_type}' + endpoint = endpoint + f"/{signal_type}" return endpoint @@ -1475,11 +1489,10 @@ class OpenTelemetry(CustomLogger): if isinstance(headers, str): # when passed HEADERS="x-honeycomb-team=B85YgLm96******" # Split only on first '=' occurrence - parts = headers.split("=", 1) - if len(parts) == 2: - _split_otel_headers = {parts[0]: parts[1]} - else: - _split_otel_headers = {} + parts = headers.split(",") + for part in parts: + key, value = part.split("=", 1) + _split_otel_headers[key] = value elif isinstance(headers, dict): _split_otel_headers = headers return _split_otel_headers diff --git a/litellm/llms/bedrock/image/amazon_titan_transformation.py b/litellm/llms/bedrock/image/amazon_titan_transformation.py new file mode 100644 index 00000000000..2709f406dfd --- /dev/null +++ b/litellm/llms/bedrock/image/amazon_titan_transformation.py @@ -0,0 +1,160 @@ +""" +Transformation logic for Amazon Titan Image Generation. +""" + +import types +from typing import List, Optional + +from openai.types.image import Image + +from litellm import get_model_info +from litellm.types.llms.bedrock import ( + AmazonNovaCanvasImageGenerationConfig, + AmazonTitanImageGenerationRequestBody, + AmazonTitanTextToImageParams, +) +from litellm.types.utils import ImageResponse + + +class AmazonTitanImageGenerationConfig: + """ + Reference: https://us-west-2.console.aws.amazon.com/bedrock/home?region=us-west-2#/providers?model=stability.stable-diffusion-xl-v0 + """ + + cfg_scale: Optional[int] = None + seed: Optional[float] = None + steps: Optional[List[str]] = None + width: Optional[int] = None + height: Optional[int] = None + + def __init__( + self, + cfg_scale: Optional[int] = None, + seed: Optional[float] = None, + steps: Optional[List[str]] = None, + width: Optional[int] = None, + height: Optional[int] = None, + ) -> None: + locals_ = locals().copy() + for key, value in locals_.items(): + if key != "self" and value is not None: + setattr(self.__class__, key, value) + + @classmethod + def get_config(cls): + return { + k: v + for k, v in cls.__dict__.items() + if not k.startswith("__") + and not isinstance( + v, + ( + types.FunctionType, + types.BuiltinFunctionType, + classmethod, + staticmethod, + ), + ) + and v is not None + } + + @classmethod + def _is_titan_model(cls, model: Optional[str] = None) -> bool: + """ + Returns True if the model is a Titan model + + Titan models follow this pattern: + + """ + if model and "amazon.titan" in model: + return True + return False + + @classmethod + def get_supported_openai_params(cls, model: Optional[str] = None) -> List: + return ["size", "n", "quality"] + + @classmethod + def map_openai_params( + cls, + non_default_params: dict, + optional_params: dict, + ): + from typing import Any, Dict + + image_generation_config: Dict[str, Any] = {} + for k, v in non_default_params.items(): + if k == "size" and v is not None: + width, height = v.split("x") + image_generation_config["width"] = int(width) + image_generation_config["height"] = int(height) + elif k == "n" and v is not None: + image_generation_config["numberOfImages"] = v + elif ( + k == "quality" and v is not None + ): # 'auto', 'hd', 'standard', 'high', 'medium', 'low' + if v in ("hd", "premium", "high"): + image_generation_config["quality"] = "premium" + elif v in ("standard", "medium", "low"): + image_generation_config["quality"] = "standard" + + if image_generation_config: + optional_params["imageGenerationConfig"] = image_generation_config + return optional_params + + @classmethod + def _transform_request( + cls, + input: str, + optional_params: dict, + ) -> AmazonTitanImageGenerationRequestBody: + from typing import Any, Dict + + image_generation_config = optional_params.pop("imageGenerationConfig", {}) + negative_text = optional_params.pop("negativeText", None) + text_to_image_params: Dict[str, Any] = {"text": input} + if negative_text: + text_to_image_params["negativeText"] = negative_text + task_type = optional_params.pop("taskType", "TEXT_IMAGE") + user_specified_image_generation_config = optional_params.pop( + "imageGenerationConfig", {} + ) + image_generation_config = { + **image_generation_config, + **user_specified_image_generation_config, + } + return AmazonTitanImageGenerationRequestBody( + taskType=task_type, + textToImageParams=AmazonTitanTextToImageParams(**text_to_image_params), # type: ignore + imageGenerationConfig=AmazonNovaCanvasImageGenerationConfig( + **image_generation_config + ), + ) + + @classmethod + def transform_response_dict_to_openai_response( + cls, model_response: ImageResponse, response_dict: dict + ) -> ImageResponse: + image_list: List[Image] = [] + for image in response_dict["images"]: + _image = Image(b64_json=image) + image_list.append(_image) + + model_response.data = image_list + + return model_response + + @classmethod + def cost_calculator( + cls, + model: str, + image_response: ImageResponse, + size: Optional[str] = None, + optional_params: Optional[dict] = None, + ) -> float: + model_info = get_model_info(model=model) + output_cost_per_image = model_info.get("output_cost_per_image") or 0.0 + if not image_response.data: + return 0.0 + num_images = len(image_response.data) + return output_cost_per_image * num_images diff --git a/litellm/llms/bedrock/image/cost_calculator.py b/litellm/llms/bedrock/image/cost_calculator.py index a0dc91d7119..9b2ae8782cb 100644 --- a/litellm/llms/bedrock/image/cost_calculator.py +++ b/litellm/llms/bedrock/image/cost_calculator.py @@ -1,6 +1,9 @@ from typing import Optional import litellm +from litellm.llms.bedrock.image.amazon_titan_transformation import ( + AmazonTitanImageGenerationConfig, +) from litellm.types.utils import ImageResponse @@ -17,6 +20,13 @@ def cost_calculator( """ if litellm.AmazonStability3Config()._is_stability_3_model(model=model): pass + elif AmazonTitanImageGenerationConfig._is_titan_model(model=model): + return AmazonTitanImageGenerationConfig.cost_calculator( + model=model, + image_response=image_response, + size=size, + optional_params=optional_params, + ) else: # Stability 1 models optional_params = optional_params or {} diff --git a/litellm/llms/bedrock/image/image_handler.py b/litellm/llms/bedrock/image/image_handler.py index 0103f190d36..1c418f78296 100644 --- a/litellm/llms/bedrock/image/image_handler.py +++ b/litellm/llms/bedrock/image/image_handler.py @@ -9,6 +9,15 @@ from pydantic import BaseModel import litellm from litellm._logging import verbose_logger from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging +from litellm.llms.bedrock.image.amazon_nova_canvas_transformation import ( + AmazonNovaCanvasConfig, +) +from litellm.llms.bedrock.image.amazon_stability3_transformation import ( + AmazonStability3Config, +) +from litellm.llms.bedrock.image.amazon_titan_transformation import ( + AmazonTitanImageGenerationConfig, +) from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, @@ -63,7 +72,7 @@ class BedrockImageGeneration(BaseAWSLLM): extra_headers=extra_headers, logging_obj=logging_obj, prompt=prompt, - api_key=api_key + api_key=api_key, ) if aimg_generation is True: @@ -190,7 +199,7 @@ class BedrockImageGeneration(BaseAWSLLM): body = json.dumps(data).encode("utf-8") headers = {"Content-Type": "application/json"} if extra_headers is not None: - headers = {"Content-Type": "application/json", **extra_headers} + headers = {"Content-Type": "application/json", **extra_headers} prepped = self.get_request_headers( credentials=boto3_credentials_info.credentials, @@ -201,7 +210,7 @@ class BedrockImageGeneration(BaseAWSLLM): headers=headers, api_key=api_key, ) - + ## LOGGING logging_obj.pre_call( input=prompt, @@ -306,15 +315,21 @@ class BedrockImageGeneration(BaseAWSLLM): if response_dict is None: raise ValueError("Error in response object format, got None") - config_class = ( - litellm.AmazonStability3Config - if litellm.AmazonStability3Config._is_stability_3_model(model=model) - else ( - litellm.AmazonNovaCanvasConfig - if litellm.AmazonNovaCanvasConfig._is_nova_model(model=model) - else litellm.AmazonStabilityConfig - ) - ) + config_class: Union[ + type[AmazonTitanImageGenerationConfig], + type[AmazonNovaCanvasConfig], + type[AmazonStability3Config], + type[litellm.AmazonStabilityConfig], + ] + if AmazonTitanImageGenerationConfig._is_titan_model(model=model): + config_class = AmazonTitanImageGenerationConfig + elif AmazonNovaCanvasConfig._is_nova_model(model=model): + config_class = AmazonNovaCanvasConfig + elif AmazonStability3Config._is_stability_3_model(model=model): + config_class = AmazonStability3Config + else: + config_class = litellm.AmazonStabilityConfig + config_class.transform_response_dict_to_openai_response( model_response=model_response, response_dict=response_dict, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 0c9a5349ab3..03992da5fe6 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -336,6 +336,24 @@ "output_cost_per_token": 0.0, "output_vector_size": 1024 }, + "amazon.titan-image-generator-v1": { + "input_cost_per_image": 0.0, + "output_cost_per_image": 0.008, + "output_cost_per_image_premium_image": 0.01, + "output_cost_per_image_above_512_and_512_pixels": 0.01, + "output_cost_per_image_above_512_and_512_pixels_and_premium_image": 0.012, + "litellm_provider": "bedrock", + "mode": "image_generation" + }, + "amazon.titan-image-generator-v2": { + "input_cost_per_image": 0.0, + "output_cost_per_image": 0.008, + "output_cost_per_image_premium_image": 0.01, + "output_cost_per_image_above_1024_and_1024_pixels": 0.01, + "output_cost_per_image_above_1024_and_1024_pixels_and_premium_image": 0.012, + "litellm_provider": "bedrock", + "mode": "image_generation" + }, "twelvelabs.marengo-embed-2-7-v1:0": { "input_cost_per_token": 7e-05, "litellm_provider": "bedrock", diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index df551c5bded..bc752dd26a0 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -1,12 +1,7 @@ import json from typing import Any, Dict, List, Literal, Optional, Union -from typing_extensions import ( - TYPE_CHECKING, - Required, - TypedDict, - override, -) +from typing_extensions import TYPE_CHECKING, Required, TypedDict, override from .openai import ChatCompletionToolCallChunk @@ -468,6 +463,15 @@ class AmazonStability3TextToImageResponse(TypedDict, total=False): finish_reasons: List[str] +class AmazonTitanTextToImageParams(TypedDict, total=False): + """ + Params for Amazon Titan Text to Image API + """ + + text: Required[str] + negativeText: str + + class AmazonNovaCanvasRequestBase(TypedDict, total=False): """ Base class for Amazon Nova Canvas API requests @@ -577,6 +581,16 @@ class AmazonNovaCanvasInpaintingRequest( imageGenerationConfig: AmazonNovaCanvasImageGenerationConfig +class AmazonTitanImageGenerationRequestBody(TypedDict, total=False): + """ + Config for Amazon Titan Image Generation API + """ + + taskType: Literal["TEXT_IMAGE", "COLOR_GUIDED_GENERATION", "INPAINTING"] + textToImageParams: AmazonTitanTextToImageParams + imageGenerationConfig: AmazonNovaCanvasImageGenerationConfig + + if TYPE_CHECKING: from botocore.awsrequest import AWSPreparedRequest else: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 0c9a5349ab3..03992da5fe6 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -336,6 +336,24 @@ "output_cost_per_token": 0.0, "output_vector_size": 1024 }, + "amazon.titan-image-generator-v1": { + "input_cost_per_image": 0.0, + "output_cost_per_image": 0.008, + "output_cost_per_image_premium_image": 0.01, + "output_cost_per_image_above_512_and_512_pixels": 0.01, + "output_cost_per_image_above_512_and_512_pixels_and_premium_image": 0.012, + "litellm_provider": "bedrock", + "mode": "image_generation" + }, + "amazon.titan-image-generator-v2": { + "input_cost_per_image": 0.0, + "output_cost_per_image": 0.008, + "output_cost_per_image_premium_image": 0.01, + "output_cost_per_image_above_1024_and_1024_pixels": 0.01, + "output_cost_per_image_above_1024_and_1024_pixels_and_premium_image": 0.012, + "litellm_provider": "bedrock", + "mode": "image_generation" + }, "twelvelabs.marengo-embed-2-7-v1:0": { "input_cost_per_token": 7e-05, "litellm_provider": "bedrock", diff --git a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py index 064fe935fd3..a2bafa85c57 100644 --- a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py +++ b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py @@ -508,3 +508,19 @@ def test_backward_compatibility_regular_nova_model(): assert result["taskType"] == "TEXT_IMAGE" assert result["textToImageParams"]["text"] == prompt assert result["imageGenerationConfig"]["cfg_scale"] == 7 + + +def test_amazon_titan_image_gen(): + from litellm import image_generation + + model_id = "bedrock/amazon.titan-image-generator-v1" + + response = litellm.image_generation( + model=model_id, + prompt="A serene mountain landscape at sunset with a lake reflection", + aws_region_name="us-east-1", + ) + + print(f"response cost: {response._hidden_params['response_cost']}") + + assert response._hidden_params["response_cost"] > 0 diff --git a/tests/llm_translation/test_bedrock_embedding.py b/tests/llm_translation/test_bedrock_embedding.py index 88918674aaa..bfdbdf53785 100644 --- a/tests/llm_translation/test_bedrock_embedding.py +++ b/tests/llm_translation/test_bedrock_embedding.py @@ -14,26 +14,41 @@ sys.path.insert( import litellm from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler -titan_embedding_response = { - "embedding": [0.1, 0.2, 0.3], - "inputTextTokenCount": 10 -} +titan_embedding_response = {"embedding": [0.1, 0.2, 0.3], "inputTextTokenCount": 10} -cohere_embedding_response = { - "embeddings": [[0.1, 0.2, 0.3]], - "inputTextTokenCount": 10 -} +cohere_embedding_response = {"embeddings": [[0.1, 0.2, 0.3]], "inputTextTokenCount": 10} img_base_64 = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAGQAAABkBAMAAACCzIhnAAAAG1BMVEURAAD///+ln5/h39/Dv79qX18uHx+If39MPz9oMSdmAAAACXBIWXMAAA7EAAAOxAGVKw4bAAABB0lEQVRYhe2SzWrEIBCAh2A0jxEs4j6GLDS9hqWmV5Flt0cJS+lRwv742DXpEjY1kOZW6HwHFZnPmVEBEARBEARB/jd0KYA/bcUYbPrRLh6amXHJ/K+ypMoyUaGthILzw0l+xI0jsO7ZcmCcm4ILd+QuVYgpHOmDmz6jBeJImdcUCmeBqQpuqRIbVmQsLCrAalrGpfoEqEogqbLTWuXCPCo+Ki1XGqgQ+jVVuhB8bOaHkvmYuzm/b0KYLWwoK58oFqi6XfxQ4Uz7d6WeKpna6ytUs5e8betMcqAv5YPC5EZB2Lm9FIn0/VP6R58+/GEY1X1egVoZ/3bt/EqF6malgSAIgiDIH+QL41409QMY0LMAAAAASUVORK5CYII=" + @pytest.mark.parametrize( "model,input_type,embed_response", [ - ("bedrock/amazon.titan-embed-text-v1", "text", titan_embedding_response), # V1 text model - ("bedrock/amazon.titan-embed-text-v2:0", "text", titan_embedding_response), # V2 text model - ("bedrock/amazon.titan-embed-image-v1", "image", titan_embedding_response), # Image model - ("bedrock/cohere.embed-english-v3", "text", cohere_embedding_response), # Cohere English - ("bedrock/cohere.embed-multilingual-v3", "text", cohere_embedding_response), # Cohere Multilingual + ( + "bedrock/amazon.titan-embed-text-v1", + "text", + titan_embedding_response, + ), # V1 text model + ( + "bedrock/amazon.titan-embed-text-v2:0", + "text", + titan_embedding_response, + ), # V2 text model + ( + "bedrock/amazon.titan-embed-image-v1", + "image", + titan_embedding_response, + ), # Image model + ( + "bedrock/cohere.embed-english-v3", + "text", + cohere_embedding_response, + ), # Cohere English + ( + "bedrock/cohere.embed-multilingual-v3", + "text", + cohere_embedding_response, + ), # Cohere Multilingual ], ) def test_bedrock_embedding_models(model, input_type, embed_response): @@ -49,7 +64,9 @@ def test_bedrock_embedding_models(model, input_type, embed_response): mock_post.return_value = mock_response # Prepare input based on type - input_data = img_base_64 if input_type == "image" else "Hello world from litellm" + input_data = ( + img_base_64 if input_type == "image" else "Hello world from litellm" + ) try: response = litellm.embedding( @@ -63,8 +80,8 @@ def test_bedrock_embedding_models(model, input_type, embed_response): # Verify response structure assert isinstance(response, litellm.EmbeddingResponse) print(response.data) - assert isinstance(response.data[0]['embedding'], list) - assert len(response.data[0]['embedding']) == 3 # Based on mock response + assert isinstance(response.data[0]["embedding"], list) + assert len(response.data[0]["embedding"]) == 3 # Based on mock response # Fetch request body request_data = json.loads(mock_post.call_args.kwargs["data"]) @@ -72,7 +89,9 @@ def test_bedrock_embedding_models(model, input_type, embed_response): # Verify AWS params are not in request body aws_params = ["aws_region_name", "aws_bedrock_runtime_endpoint"] for param in aws_params: - assert param not in request_data, f"AWS param {param} should not be in request body" + assert ( + param not in request_data + ), f"AWS param {param} should not be in request body" except Exception as e: pytest.fail(f"Error occurred: {e}") @@ -85,7 +104,7 @@ def test_e2e_bedrock_embedding(): """ print("Testing text embedding...") original_region_name = os.environ.get("AWS_REGION_NAME") - + os.environ["AWS_REGION_NAME"] = "us-east-1" litellm._turn_on_debug() @@ -93,34 +112,47 @@ def test_e2e_bedrock_embedding(): model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0", input=["Hello world from LiteLLM with TwelveLabs Marengo!"], ) - + # Validate response structure - assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type" - assert hasattr(response, 'data'), "Response should have 'data' attribute" + assert isinstance( + response, litellm.EmbeddingResponse + ), "Response should be EmbeddingResponse type" + assert hasattr(response, "data"), "Response should have 'data' attribute" assert len(response.data) > 0, "Response data should not be empty" - + # Validate first embedding embedding_obj = response.data[0] - assert hasattr(embedding_obj, 'embedding'), "Embedding object should have 'embedding' attribute" - assert isinstance(embedding_obj.embedding, list), "Embedding should be a list of floats" + assert hasattr( + embedding_obj, "embedding" + ), "Embedding object should have 'embedding' attribute" + assert isinstance( + embedding_obj.embedding, list + ), "Embedding should be a list of floats" assert len(embedding_obj.embedding) > 0, "Embedding vector should not be empty" - assert all(isinstance(x, (int, float)) for x in embedding_obj.embedding), "All embedding values should be numeric" - + assert all( + isinstance(x, (int, float)) for x in embedding_obj.embedding + ), "All embedding values should be numeric" + # Validate embedding properties assert embedding_obj.index == 0, "First embedding should have index 0" - assert embedding_obj.object == "embedding", "Embedding object type should be 'embedding'" - + assert ( + embedding_obj.object == "embedding" + ), "Embedding object type should be 'embedding'" + # Validate usage information - assert hasattr(response, 'usage'), "Response should have usage information" + assert hasattr(response, "usage"), "Response should have usage information" assert response.usage is not None, "Usage should not be None" assert response.usage.total_tokens >= 0, "Total tokens should be non-negative" - - print(f"Text embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}") + + print( + f"Text embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}" + ) # Restore original region name if original_region_name: os.environ["AWS_REGION_NAME"] = original_region_name + def test_e2e_bedrock_embedding_image_twelvelabs_marengo(): """ Test image embedding with TwelveLabs Marengo. @@ -130,46 +162,60 @@ def test_e2e_bedrock_embedding_image_twelvelabs_marengo(): original_region_name = os.environ.get("AWS_REGION_NAME") os.environ["AWS_REGION_NAME"] = "us-east-1" litellm._turn_on_debug() - + # Load duck.png and convert to base64 duck_img_path = os.path.join(os.path.dirname(__file__), "duck.png") with open(duck_img_path, "rb") as img_file: - duck_img_data = base64.b64encode(img_file.read()).decode('utf-8') + duck_img_data = base64.b64encode(img_file.read()).decode("utf-8") duck_img_base64 = f"data:image/png;base64,{duck_img_data}" - + response = litellm.embedding( model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0", input=[duck_img_base64], aws_region_name="us-east-1", - input_type="image" + input_type="image", ) - + # Validate response structure - assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type" - assert hasattr(response, 'data'), "Response should have 'data' attribute" + assert isinstance( + response, litellm.EmbeddingResponse + ), "Response should be EmbeddingResponse type" + assert hasattr(response, "data"), "Response should have 'data' attribute" assert len(response.data) > 0, "Response data should not be empty" - + # Validate first embedding embedding_obj = response.data[0] - assert hasattr(embedding_obj, 'embedding'), "Embedding object should have 'embedding' attribute" - assert isinstance(embedding_obj.embedding, list), "Embedding should be a list of floats" + assert hasattr( + embedding_obj, "embedding" + ), "Embedding object should have 'embedding' attribute" + assert isinstance( + embedding_obj.embedding, list + ), "Embedding should be a list of floats" assert len(embedding_obj.embedding) > 0, "Embedding vector should not be empty" - assert all(isinstance(x, (int, float)) for x in embedding_obj.embedding), "All embedding values should be numeric" - + assert all( + isinstance(x, (int, float)) for x in embedding_obj.embedding + ), "All embedding values should be numeric" + # Validate embedding properties assert embedding_obj.index == 0, "First embedding should have index 0" - assert embedding_obj.object == "embedding", "Embedding object type should be 'embedding'" - + assert ( + embedding_obj.object == "embedding" + ), "Embedding object type should be 'embedding'" + # Validate usage information - assert hasattr(response, 'usage'), "Response should have usage information" + assert hasattr(response, "usage"), "Response should have usage information" assert response.usage is not None, "Usage should not be None" assert response.usage.total_tokens >= 0, "Total tokens should be non-negative" - + # TwelveLabs Marengo should return 1024-dimensional embeddings expected_dimension = 1024 - assert len(embedding_obj.embedding) == expected_dimension, f"TwelveLabs Marengo should return {expected_dimension}-dimensional embeddings, got {len(embedding_obj.embedding)}" - - print(f"Image embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}") + assert ( + len(embedding_obj.embedding) == expected_dimension + ), f"TwelveLabs Marengo should return {expected_dimension}-dimensional embeddings, got {len(embedding_obj.embedding)}" + + print( + f"Image embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}" + ) # Restore original region name if original_region_name: @@ -185,37 +231,54 @@ def test_e2e_bedrock_async_invoke_embedding_twelvelabs_marengo(): original_region_name = os.environ.get("AWS_REGION_NAME") os.environ["AWS_REGION_NAME"] = "us-east-1" litellm._turn_on_debug() - + # Mock the HTTP call to return async invoke response - with patch("litellm.llms.bedrock.embed.embedding.BedrockEmbedding._make_sync_call") as mock_call: + with patch( + "litellm.llms.bedrock.embed.embedding.BedrockEmbedding._make_sync_call" + ) as mock_call: mock_call.return_value = { "invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-job-123" } - + response = litellm.embedding( model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", input=["Hello world from LiteLLM async invoke!"], aws_region_name="us-east-1", inputType="text", - output_s3_uri="s3://test-bucket/async-invoke-output/" + output_s3_uri="s3://test-bucket/async-invoke-output/", ) - + # Validate response structure - assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type" - assert hasattr(response, '_hidden_params'), "Response should have _hidden_params" + assert isinstance( + response, litellm.EmbeddingResponse + ), "Response should be EmbeddingResponse type" + assert hasattr( + response, "_hidden_params" + ), "Response should have _hidden_params" assert response._hidden_params is not None, "Hidden params should not be None" - + # Validate hidden params contain invocation ARN - assert hasattr(response._hidden_params, '_invocation_arn'), "Hidden params should have _invocation_arn" - assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-job-123", "Invocation ARN should be preserved" - + assert hasattr( + response._hidden_params, "_invocation_arn" + ), "Hidden params should have _invocation_arn" + assert ( + response._hidden_params._invocation_arn + == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-job-123" + ), "Invocation ARN should be preserved" + # Validate embedding structure assert len(response.data) == 1, "Should have one embedding" - assert response.data[0].object == "embedding", "Embedding object should be 'embedding'" - assert response.data[0].embedding == [], "Embedding should be empty for async jobs" - - print(f"Async invoke embedding successful! Invocation ARN: {response._hidden_params._invocation_arn}") - + assert ( + response.data[0].object == "embedding" + ), "Embedding object should be 'embedding'" + assert ( + response.data[0].embedding == [] + ), "Embedding should be empty for async jobs" + + print( + f"Async invoke embedding successful! Invocation ARN: {response._hidden_params._invocation_arn}" + ) + # Restore original region name if original_region_name: os.environ["AWS_REGION_NAME"] = original_region_name @@ -231,32 +294,45 @@ async def test_e2e_bedrock_async_invoke_embedding_async_twelvelabs_marengo(): original_region_name = os.environ.get("AWS_REGION_NAME") os.environ["AWS_REGION_NAME"] = "us-east-1" litellm._turn_on_debug() - + # Mock the async HTTP call to return async invoke response - with patch("litellm.llms.bedrock.embed.embedding.BedrockEmbedding._make_async_call") as mock_call: + with patch( + "litellm.llms.bedrock.embed.embedding.BedrockEmbedding._make_async_call" + ) as mock_call: mock_call.return_value = { "invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-async-job-456" } - + response = await litellm.aembedding( model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", input=["Hello world from LiteLLM async invoke async!"], aws_region_name="us-east-1", inputType="text", - output_s3_uri="s3://test-bucket/async-invoke-output/" + output_s3_uri="s3://test-bucket/async-invoke-output/", ) - + # Validate response structure - assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type" - assert hasattr(response, '_hidden_params'), "Response should have _hidden_params" + assert isinstance( + response, litellm.EmbeddingResponse + ), "Response should be EmbeddingResponse type" + assert hasattr( + response, "_hidden_params" + ), "Response should have _hidden_params" assert response._hidden_params is not None, "Hidden params should not be None" - + # Validate hidden params contain invocation ARN - assert hasattr(response._hidden_params, '_invocation_arn'), "Hidden params should have _invocation_arn" - assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-async-job-456", "Invocation ARN should be preserved" - - print(f"Async invoke embedding successful! Invocation ARN: {response._hidden_params._invocation_arn}") - + assert hasattr( + response._hidden_params, "_invocation_arn" + ), "Hidden params should have _invocation_arn" + assert ( + response._hidden_params._invocation_arn + == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-async-job-456" + ), "Invocation ARN should be preserved" + + print( + f"Async invoke embedding successful! Invocation ARN: {response._hidden_params._invocation_arn}" + ) + # Restore original region name if original_region_name: - os.environ["AWS_REGION_NAME"] = original_region_name \ No newline at end of file + os.environ["AWS_REGION_NAME"] = original_region_name diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index fc33b736872..496136b3aea 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -1,24 +1,24 @@ import json import os import sys -import unittest -from unittest.mock import MagicMock, patch -from datetime import datetime, timedelta import time +import unittest +from datetime import datetime, timedelta +from unittest.mock import MagicMock, patch # Adds the grandparent directory to sys.path to allow importing project modules sys.path.insert(0, os.path.abspath("../..")) -from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig -from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - -from opentelemetry.sdk.trace import TracerProvider -from opentelemetry.sdk.trace.export import SimpleSpanProcessor from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider -from opentelemetry.sdk._logs.export import SimpleLogRecordProcessor, InMemoryLogExporter +from opentelemetry.sdk._logs.export import InMemoryLogExporter, SimpleLogRecordProcessor from opentelemetry.sdk.metrics import MeterProvider from opentelemetry.sdk.metrics.export import InMemoryMetricReader +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps + class TestOpenTelemetryGuardrails(unittest.TestCase): @patch("litellm.integrations.opentelemetry.datetime") @@ -604,19 +604,35 @@ class TestOpenTelemetry(unittest.TestCase): logs = self.wait_for_log(log_exporter, "gen_ai.") self.assertTrue(logs, "Expected at least one gen_ai log") - user_logs = [log for log in logs if log.log_record.attributes.get("event_name") == "gen_ai.content.prompt"] + user_logs = [ + log + for log in logs + if log.log_record.attributes.get("event_name") == "gen_ai.content.prompt" + ] self.assertTrue(user_logs, "did not see a gen_ai.content.prompt log") # check log bodies user_prompt = user_logs[0].log_record.attributes.get("gen_ai.prompt") - self.assertEqual("What is the capital of France?", user_prompt, "did not see a prompt message") + self.assertEqual( + "What is the capital of France?", + user_prompt, + "did not see a prompt message", + ) - choice_logs = [log for log in logs if log.log_record.attributes.get("event_name") == "gen_ai.content.completion"] + choice_logs = [ + log + for log in logs + if log.log_record.attributes.get("event_name") + == "gen_ai.content.completion" + ] self.assertTrue(choice_logs, "did not see a gen_ai.content.completion event") choice_response = choice_logs[0].log_record.body self.assertIsNotNone(choice_response, "did not see a response message") - self.assertEqual("stop", choice_response.get("finish_reason"), "did not see expected finish reason") - + self.assertEqual( + "stop", + choice_response.get("finish_reason"), + "did not see expected finish reason", + ) def test_handle_success_spans_only(self): # make sure neither events nor metrics is on @@ -670,8 +686,7 @@ class TestOpenTelemetry(unittest.TestCase): # ) # model attribute should be on that span found = any( - s.attributes - and s.attributes.get("gen_ai.request.model") == self.MODEL + s.attributes and s.attributes.get("gen_ai.request.model") == self.MODEL for s in spans ) self.assertTrue(found, "expected gen_ai.request.model on span attributes") @@ -755,13 +770,7 @@ class TestOpenTelemetry(unittest.TestCase): def test_get_span_name_with_generation_name(self): """Test _get_span_name returns generation_name when present""" otel = OpenTelemetry() - kwargs = { - "litellm_params": { - "metadata": { - "generation_name": "custom_span" - } - } - } + kwargs = {"litellm_params": {"metadata": {"generation_name": "custom_span"}}} result = otel._get_span_name(kwargs) self.assertEqual(result, "custom_span") @@ -774,7 +783,7 @@ class TestOpenTelemetry(unittest.TestCase): result = otel._get_span_name(kwargs) self.assertEqual(result, LITELLM_REQUEST_SPAN_NAME) - @patch('litellm.turn_off_message_logging', False) + @patch("litellm.turn_off_message_logging", False) def test_maybe_log_raw_request_creates_span(self): """Test _maybe_log_raw_request creates span when logging enabled""" from litellm.integrations.opentelemetry import RAW_REQUEST_SPAN_NAME @@ -790,12 +799,16 @@ class TestOpenTelemetry(unittest.TestCase): otel._to_ns = MagicMock(return_value=1234567890) kwargs = {"litellm_params": {"metadata": {}}} - otel._maybe_log_raw_request(kwargs, {}, datetime.now(), datetime.now(), MagicMock()) + otel._maybe_log_raw_request( + kwargs, {}, datetime.now(), datetime.now(), MagicMock() + ) mock_tracer.start_span.assert_called_once() - self.assertEqual(mock_tracer.start_span.call_args[1]['name'], RAW_REQUEST_SPAN_NAME) + self.assertEqual( + mock_tracer.start_span.call_args[1]["name"], RAW_REQUEST_SPAN_NAME + ) - @patch('litellm.turn_off_message_logging', True) + @patch("litellm.turn_off_message_logging", True) def test_maybe_log_raw_request_skips_when_logging_disabled(self): """Test _maybe_log_raw_request skips when logging disabled""" otel = OpenTelemetry() @@ -803,24 +816,50 @@ class TestOpenTelemetry(unittest.TestCase): otel.get_tracer_to_use_for_request = MagicMock(return_value=mock_tracer) kwargs = {"litellm_params": {"metadata": {}}} - otel._maybe_log_raw_request(kwargs, {}, datetime.now(), datetime.now(), MagicMock()) + otel._maybe_log_raw_request( + kwargs, {}, datetime.now(), datetime.now(), MagicMock() + ) mock_tracer.start_span.assert_not_called() +class TestOpenTelemetryHeaderSplitting(unittest.TestCase): + """Test suite for _get_headers_dictionary method""" + + def test_split_multiple_headers_comma_separated(self): + """Test splitting multiple headers separated by commas""" + otel = OpenTelemetry() + headers = "api-key=key,other-config-value=value" + result = otel._get_headers_dictionary(headers) + self.assertEqual(result, {"api-key": "key", "other-config-value": "value"}) + + def test_split_headers_with_equals_in_values(self): + """Test splitting headers where values contain equals signs (split only on first '=')""" + otel = OpenTelemetry() + headers = "api-key=value1=part2,config=setting=enabled" + result = otel._get_headers_dictionary(headers) + self.assertEqual( + result, {"api-key": "value1=part2", "config": "setting=enabled"} + ) + + class TestOpenTelemetryEndpointNormalization(unittest.TestCase): """Test suite for the unified _normalize_otel_endpoint method""" def test_normalize_traces_endpoint_from_logs_path(self): """Test normalizing endpoint with /v1/logs to /v1/traces""" otel = OpenTelemetry() - result = otel._normalize_otel_endpoint("http://collector:4318/v1/logs", "traces") + result = otel._normalize_otel_endpoint( + "http://collector:4318/v1/logs", "traces" + ) self.assertEqual(result, "http://collector:4318/v1/traces") def test_normalize_traces_endpoint_from_metrics_path(self): """Test normalizing endpoint with /v1/metrics to /v1/traces""" otel = OpenTelemetry() - result = otel._normalize_otel_endpoint("http://collector:4318/v1/metrics", "traces") + result = otel._normalize_otel_endpoint( + "http://collector:4318/v1/metrics", "traces" + ) self.assertEqual(result, "http://collector:4318/v1/traces") def test_normalize_traces_endpoint_from_base_url(self): @@ -838,19 +877,25 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase): def test_normalize_traces_endpoint_already_correct(self): """Test endpoint already ending with /v1/traces remains unchanged""" otel = OpenTelemetry() - result = otel._normalize_otel_endpoint("http://collector:4318/v1/traces", "traces") + result = otel._normalize_otel_endpoint( + "http://collector:4318/v1/traces", "traces" + ) self.assertEqual(result, "http://collector:4318/v1/traces") def test_normalize_metrics_endpoint_from_traces_path(self): """Test normalizing endpoint with /v1/traces to /v1/metrics""" otel = OpenTelemetry() - result = otel._normalize_otel_endpoint("http://collector:4318/v1/traces", "metrics") + result = otel._normalize_otel_endpoint( + "http://collector:4318/v1/traces", "metrics" + ) self.assertEqual(result, "http://collector:4318/v1/metrics") def test_normalize_metrics_endpoint_from_logs_path(self): """Test normalizing endpoint with /v1/logs to /v1/metrics""" otel = OpenTelemetry() - result = otel._normalize_otel_endpoint("http://collector:4318/v1/logs", "metrics") + result = otel._normalize_otel_endpoint( + "http://collector:4318/v1/logs", "metrics" + ) self.assertEqual(result, "http://collector:4318/v1/metrics") def test_normalize_metrics_endpoint_from_base_url(self): @@ -862,19 +907,25 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase): def test_normalize_metrics_endpoint_already_correct(self): """Test endpoint already ending with /v1/metrics remains unchanged""" otel = OpenTelemetry() - result = otel._normalize_otel_endpoint("http://collector:4318/v1/metrics", "metrics") + result = otel._normalize_otel_endpoint( + "http://collector:4318/v1/metrics", "metrics" + ) self.assertEqual(result, "http://collector:4318/v1/metrics") def test_normalize_logs_endpoint_from_traces_path(self): """Test normalizing endpoint with /v1/traces to /v1/logs""" otel = OpenTelemetry() - result = otel._normalize_otel_endpoint("http://collector:4318/v1/traces", "logs") + result = otel._normalize_otel_endpoint( + "http://collector:4318/v1/traces", "logs" + ) self.assertEqual(result, "http://collector:4318/v1/logs") def test_normalize_logs_endpoint_from_metrics_path(self): """Test normalizing endpoint with /v1/metrics to /v1/logs""" otel = OpenTelemetry() - result = otel._normalize_otel_endpoint("http://collector:4318/v1/metrics", "logs") + result = otel._normalize_otel_endpoint( + "http://collector:4318/v1/metrics", "logs" + ) self.assertEqual(result, "http://collector:4318/v1/logs") def test_normalize_logs_endpoint_from_base_url(self): @@ -912,7 +963,7 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase): otel = OpenTelemetry() endpoint = "http://collector:4318/v1/traces" - with patch('litellm._logging.verbose_logger.warning') as mock_warning: + with patch("litellm._logging.verbose_logger.warning") as mock_warning: result = otel._normalize_otel_endpoint(endpoint, "invalid") # Should return endpoint unchanged @@ -924,18 +975,24 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase): call_args = mock_warning.call_args[0] self.assertIn("Invalid signal_type", call_args[0]) self.assertEqual(call_args[1], "invalid") # signal_type parameter - self.assertEqual(call_args[2], {'traces', 'metrics', 'logs'}) # valid_signals parameter + self.assertEqual( + call_args[2], {"traces", "metrics", "logs"} + ) # valid_signals parameter def test_normalize_endpoint_https(self): """Test normalization works with https URLs""" otel = OpenTelemetry() - result = otel._normalize_otel_endpoint("https://collector.example.com:4318", "logs") + result = otel._normalize_otel_endpoint( + "https://collector.example.com:4318", "logs" + ) self.assertEqual(result, "https://collector.example.com:4318/v1/logs") def test_normalize_endpoint_with_path_prefix(self): """Test normalization works with URLs that have path prefixes""" otel = OpenTelemetry() - result = otel._normalize_otel_endpoint("http://collector:4318/otel/v1/traces", "logs") + result = otel._normalize_otel_endpoint( + "http://collector:4318/otel/v1/traces", "logs" + ) # Should replace the final /traces with /logs self.assertEqual(result, "http://collector:4318/otel/v1/logs") @@ -984,8 +1041,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): from opentelemetry.sdk.trace.export import BatchSpanProcessor config = OpenTelemetryConfig( - exporter="otlp_http", - endpoint="http://collector:4318" + exporter="otlp_http", endpoint="http://collector:4318" ) otel = OpenTelemetry(config=config) @@ -1005,8 +1061,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): from opentelemetry.sdk.trace.export import BatchSpanProcessor config = OpenTelemetryConfig( - exporter="otlp_grpc", - endpoint="http://collector:4317" + exporter="otlp_grpc", endpoint="http://collector:4317" ) otel = OpenTelemetry(config=config) @@ -1025,10 +1080,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): ) from opentelemetry.sdk.trace.export import BatchSpanProcessor - config = OpenTelemetryConfig( - exporter="grpc", - endpoint="http://collector:4317" - ) + config = OpenTelemetryConfig(exporter="grpc", endpoint="http://collector:4317") otel = OpenTelemetry(config=config) processor = otel._get_span_processor() @@ -1047,8 +1099,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): from opentelemetry.sdk.trace.export import BatchSpanProcessor config = OpenTelemetryConfig( - exporter="http/protobuf", - endpoint="http://collector:4318" + exporter="http/protobuf", endpoint="http://collector:4318" ) otel = OpenTelemetry(config=config) @@ -1083,9 +1134,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): from opentelemetry.exporter.otlp.proto.http._log_exporter import OTLPLogExporter config = OpenTelemetryConfig( - exporter="otlp_http", - endpoint="http://collector:4318", - enable_events=True + exporter="otlp_http", endpoint="http://collector:4318", enable_events=True ) otel = OpenTelemetry(config=config) @@ -1095,16 +1144,14 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): self.assertIsInstance(exporter, OTLPLogExporter) # Check that it's from the http module by checking the module name - self.assertIn('http', exporter.__class__.__module__) + self.assertIn("http", exporter.__class__.__module__) def test_get_log_exporter_uses_grpc_exporter_for_otlp_grpc(self): """Test that otlp_grpc protocol uses gRPC OTLPLogExporter""" from opentelemetry.exporter.otlp.proto.grpc._log_exporter import OTLPLogExporter config = OpenTelemetryConfig( - exporter="otlp_grpc", - endpoint="http://collector:4317", - enable_events=True + exporter="otlp_grpc", endpoint="http://collector:4317", enable_events=True ) otel = OpenTelemetry(config=config) @@ -1114,16 +1161,14 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): self.assertIsInstance(exporter, OTLPLogExporter) # Check that it's from the grpc module by checking the module name - self.assertIn('grpc', exporter.__class__.__module__) + self.assertIn("grpc", exporter.__class__.__module__) def test_get_log_exporter_uses_grpc_exporter_for_grpc_alias(self): """Test that 'grpc' protocol alias uses gRPC OTLPLogExporter""" from opentelemetry.exporter.otlp.proto.grpc._log_exporter import OTLPLogExporter config = OpenTelemetryConfig( - exporter="grpc", - endpoint="http://collector:4317", - enable_events=True + exporter="grpc", endpoint="http://collector:4317", enable_events=True ) otel = OpenTelemetry(config=config) @@ -1133,16 +1178,13 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): self.assertIsInstance(exporter, OTLPLogExporter) # Check that it's from the grpc module by checking the module name - self.assertIn('grpc', exporter.__class__.__module__) + self.assertIn("grpc", exporter.__class__.__module__) def test_get_log_exporter_uses_console_exporter_for_console(self): """Test that console protocol uses ConsoleLogExporter""" from opentelemetry.sdk._logs.export import ConsoleLogExporter - config = OpenTelemetryConfig( - exporter="console", - enable_events=True - ) + config = OpenTelemetryConfig(exporter="console", enable_events=True) otel = OpenTelemetry(config=config) exporter = otel._get_log_exporter() @@ -1154,13 +1196,10 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): """Test that unknown protocol defaults to ConsoleLogExporter with warning""" from opentelemetry.sdk._logs.export import ConsoleLogExporter - config = OpenTelemetryConfig( - exporter="unknown_protocol", - enable_events=True - ) + config = OpenTelemetryConfig(exporter="unknown_protocol", enable_events=True) otel = OpenTelemetry(config=config) - with patch('litellm._logging.verbose_logger.warning') as mock_warning: + with patch("litellm._logging.verbose_logger.warning") as mock_warning: exporter = otel._get_log_exporter() # Verify the exporter defaults to console @@ -1172,7 +1211,14 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): self.assertIn("Unknown log exporter", args[0]) self.assertIn("unknown_protocol", args[1]) - @patch.dict(os.environ, {"OTEL_EXPORTER": "otlp_http", "OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4318"}, clear=False) + @patch.dict( + os.environ, + { + "OTEL_EXPORTER": "otlp_http", + "OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4318", + }, + clear=False, + ) def test_protocol_selection_from_environment_http(self): """Test that protocol selection works correctly from environment variables for HTTP""" from opentelemetry.exporter.otlp.proto.http.trace_exporter import ( @@ -1189,7 +1235,14 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): self.assertIsInstance(processor, BatchSpanProcessor) self.assertIsInstance(processor.span_exporter, OTLPSpanExporterHTTP) - @patch.dict(os.environ, {"OTEL_EXPORTER": "otlp_grpc", "OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4317"}, clear=False) + @patch.dict( + os.environ, + { + "OTEL_EXPORTER": "otlp_grpc", + "OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4317", + }, + clear=False, + ) def test_protocol_selection_from_environment_grpc(self): """Test that protocol selection works correctly from environment variables for gRPC""" from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import ( @@ -1209,8 +1262,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): def test_http_exporter_endpoint_normalization_for_traces(self): """Test that HTTP trace exporter gets properly normalized endpoint""" config = OpenTelemetryConfig( - exporter="otlp_http", - endpoint="http://collector:4318" + exporter="otlp_http", endpoint="http://collector:4318" ) otel = OpenTelemetry(config=config) @@ -1218,14 +1270,13 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): # Verify the endpoint was normalized to include /v1/traces # Access the private _endpoint attribute if available - if hasattr(processor.span_exporter, '_endpoint'): + if hasattr(processor.span_exporter, "_endpoint"): self.assertEqual(processor.span_exporter._endpoint, "http://collector:4318/v1/traces") # type: ignore[attr-defined] def test_grpc_exporter_endpoint_normalization_for_traces(self): """Test that gRPC trace exporter gets properly normalized endpoint""" config = OpenTelemetryConfig( - exporter="otlp_grpc", - endpoint="http://collector:4317" + exporter="otlp_grpc", endpoint="http://collector:4317" ) otel = OpenTelemetry(config=config) @@ -1233,12 +1284,14 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): # Verify the endpoint was normalized to include /v1/traces # Note: gRPC exporters strip the http:// prefix, so we check for the normalized path - if hasattr(processor.span_exporter, '_endpoint'): + if hasattr(processor.span_exporter, "_endpoint"): # gRPC exporter strips http:// prefix - self.assertIn('collector:4317', processor.span_exporter._endpoint) # type: ignore[attr-defined] + self.assertIn("collector:4317", processor.span_exporter._endpoint) # type: ignore[attr-defined] # The endpoint should have been normalized with /v1/traces before being passed to gRPC exporter # We verify this by checking the normalization function was called correctly - normalized = otel._normalize_otel_endpoint("http://collector:4317", "traces") + normalized = otel._normalize_otel_endpoint( + "http://collector:4317", "traces" + ) self.assertEqual(normalized, "http://collector:4317/v1/traces") def test_http_log_exporter_endpoint_normalization_for_logs(self): @@ -1246,7 +1299,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): config = OpenTelemetryConfig( exporter="otlp_http", endpoint="http://collector:4318/v1/traces", - enable_events=True + enable_events=True, ) otel = OpenTelemetry(config=config) @@ -1254,7 +1307,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): # Verify the endpoint was normalized to /v1/logs (not /v1/traces) # Access the private _endpoint attribute if available - if hasattr(exporter, '_endpoint'): + if hasattr(exporter, "_endpoint"): self.assertEqual(exporter._endpoint, "http://collector:4318/v1/logs") # type: ignore[attr-defined] def test_grpc_log_exporter_endpoint_normalization_for_logs(self): @@ -1262,7 +1315,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): config = OpenTelemetryConfig( exporter="otlp_grpc", endpoint="http://collector:4317/v1/traces", - enable_events=True + enable_events=True, ) otel = OpenTelemetry(config=config) @@ -1270,10 +1323,12 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): # Verify the endpoint was normalized to /v1/logs (not /v1/traces) # Note: gRPC exporters strip the http:// prefix, so we check for the normalized path - if hasattr(exporter, '_endpoint'): + if hasattr(exporter, "_endpoint"): # gRPC exporter strips http:// prefix - self.assertIn('collector:4317', exporter._endpoint) # type: ignore[attr-defined] + self.assertIn("collector:4317", exporter._endpoint) # type: ignore[attr-defined] # The endpoint should have been normalized with /v1/logs before being passed to gRPC exporter # We verify this by checking the normalization function was called correctly - normalized = otel._normalize_otel_endpoint("http://collector:4317/v1/traces", "logs") + normalized = otel._normalize_otel_endpoint( + "http://collector:4317/v1/traces", "logs" + ) self.assertEqual(normalized, "http://collector:4317/v1/logs")