mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(opentelemetry.py): fix issue where headers were not being split correctly + feat(bedrock): add titan image generations w/ cost tracking (#15916)
* fix(opentelemetry.py): fix issue where headers were not being split correctly * feat(bedrock/image): Support bedrock titan image generation Closes https://github.com/BerriAI/litellm/issues/361 * build(model_prices_and_context_window.json): track titan image gen pricing enables cost tracking per request * feat(amazon_titan_transformation.py): support titan image generation cost tracking * docs: document new model * docs: update docs to indicate cost tracking + refactor rerank into separate doc * fix: fix mypy linting error * fix: fix type ignore
This commit is contained in:
parent
72bbdfd3f3
commit
346e036399
16 changed files with 859 additions and 439 deletions
|
|
@ -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
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```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}")
|
||||
```
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
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"
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### 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:
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```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}")
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
```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"
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## 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
|
||||
|
||||
<Tabs>
|
||||
<TabItem label="SDK" value="sdk">
|
||||
|
||||
```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)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem label="PROXY" value="proxy">
|
||||
|
||||
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
|
||||
|
||||
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
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:
|
|||
|
||||
</Tabs>
|
||||
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
150
docs/my-website/docs/providers/bedrock_image_gen.md
Normal file
150
docs/my-website/docs/providers/bedrock_image_gen.md
Normal file
|
|
@ -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
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
### 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}")
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
### 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"]}
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## 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:
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```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}")
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY">
|
||||
|
||||
```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"
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Authentication
|
||||
|
||||
All standard Bedrock authentication methods are supported for image generation. See [Bedrock Authentication](./bedrock#boto3---authentication) for details.
|
||||
|
||||
94
docs/my-website/docs/providers/bedrock_rerank.md
Normal file
94
docs/my-website/docs/providers/bedrock_rerank.md
Normal file
|
|
@ -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
|
||||
|
||||
<Tabs>
|
||||
<TabItem label="SDK" value="sdk">
|
||||
|
||||
```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)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem label="PROXY" value="proxy">
|
||||
|
||||
### 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
|
||||
|
||||
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Authentication
|
||||
|
||||
All standard Bedrock authentication methods are supported for rerank. See [Bedrock Authentication](./bedrock#boto3---authentication) for details.
|
||||
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
160
litellm/llms/bedrock/image/amazon_titan_transformation.py
Normal file
160
litellm/llms/bedrock/image/amazon_titan_transformation.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
os.environ["AWS_REGION_NAME"] = original_region_name
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue