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")