mirror of
https://github.com/tiennm99/litellm.git
synced 2026-07-25 04:20:23 +00:00
feat: add sagemaker_nova provider for Amazon Nova models on SageMaker (#21542)
* feat: add sagemaker_nova provider for Nova models on SageMaker Add support for custom/fine-tuned Amazon Nova models (Nova Micro, Nova Lite, Nova 2 Lite) deployed on SageMaker Inference real-time endpoints. Nova uses OpenAI-compatible request/response format with additional Nova-specific parameters (top_k, reasoning_effort, allowed_token_ids, truncate_prompt_tokens) and requires stream:true in the request body. Nova endpoints also reject 'model' in the request body. Changes: - New provider: sagemaker_nova/<endpoint-name> - SagemakerNovaConfig inherits from SagemakerChatConfig - Override transform_request to strip 'model' from request body - Override supports_stream_param_in_request_body (True for Nova) - Extend get_supported_openai_params with Nova-specific params - Refactored SagemakerChatConfig to use custom_llm_provider param instead of hardcoded strings (backwards-compatible) - Consolidated main.py routing for sagemaker_chat and sagemaker_nova - 22 unit tests + 9 integration tests (skip-gated) - Documentation with SDK, streaming, multimodal, and proxy examples - All tests verified against live SageMaker Nova endpoint * fix: move integration tests to tests/local_testing/ per test directory policy * fix: remove unused module-level SagemakerNovaConfig instance The sagemaker_nova_config singleton was never imported or used — the ProviderConfigManager creates its own instance via the lambda registered in utils.py. Removing this leftover boilerplate. --------- Co-authored-by: Krish Dholakia <krrishdholakia@gmail.com>
This commit is contained in:
co-authored by
Krish Dholakia
parent
a2708f332f
commit
374c3458d5
@@ -526,3 +526,98 @@ print(f"response: {response}")
|
||||
```
|
||||
|
||||
|
||||
|
||||
## Nova Models on SageMaker
|
||||
|
||||
LiteLLM supports Amazon Nova models (Nova Micro, Nova Lite, Nova 2 Lite) deployed on SageMaker Inference real-time endpoints. These custom/fine-tuned Nova models use an OpenAI-compatible API format.
|
||||
|
||||
**Reference:** [AWS Blog - Amazon SageMaker Inference for Custom Amazon Nova Models](https://aws.amazon.com/blogs/aws/announcing-amazon-sagemaker-inference-for-custom-amazon-nova-models/)
|
||||
|
||||
### Usage
|
||||
|
||||
Use the `sagemaker_nova/` prefix with your SageMaker endpoint name:
|
||||
|
||||
```python
|
||||
import litellm
|
||||
import os
|
||||
|
||||
os.environ["AWS_ACCESS_KEY_ID"] = ""
|
||||
os.environ["AWS_SECRET_ACCESS_KEY"] = ""
|
||||
os.environ["AWS_REGION_NAME"] = "us-east-1"
|
||||
|
||||
# Basic chat completion
|
||||
response = litellm.completion(
|
||||
model="sagemaker_nova/my-nova-endpoint",
|
||||
messages=[{"role": "user", "content": "Hello, how are you?"}],
|
||||
temperature=0.7,
|
||||
max_tokens=512,
|
||||
)
|
||||
print(response.choices[0].message.content)
|
||||
```
|
||||
|
||||
### Streaming
|
||||
|
||||
```python
|
||||
response = litellm.completion(
|
||||
model="sagemaker_nova/my-nova-endpoint",
|
||||
messages=[{"role": "user", "content": "Write a short poem"}],
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
for chunk in response:
|
||||
if chunk.choices[0].delta.content:
|
||||
print(chunk.choices[0].delta.content, end="")
|
||||
```
|
||||
|
||||
### Multimodal (Images)
|
||||
|
||||
Nova models on SageMaker support image inputs using base64 data URIs:
|
||||
|
||||
```python
|
||||
response = litellm.completion(
|
||||
model="sagemaker_nova/my-nova-endpoint",
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "What's in this image?"},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/jpeg;base64,..."}}
|
||||
]
|
||||
}
|
||||
],
|
||||
)
|
||||
```
|
||||
|
||||
### Proxy Config
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: nova-micro
|
||||
litellm_params:
|
||||
model: sagemaker_nova/my-nova-micro-endpoint
|
||||
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
|
||||
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
|
||||
aws_region_name: us-east-1
|
||||
```
|
||||
|
||||
### Supported Parameters
|
||||
|
||||
All standard OpenAI parameters are supported, plus these Nova-specific parameters:
|
||||
|
||||
| Parameter | Type | Description |
|
||||
|-----------|------|-------------|
|
||||
| `top_k` | integer | Limits token selection to top K most likely tokens |
|
||||
| `reasoning_effort` | `"low"` \| `"high"` | Reasoning effort level (Nova 2 Lite custom models only) |
|
||||
| `allowed_token_ids` | array[int] | Restrict output to specified token IDs |
|
||||
| `truncate_prompt_tokens` | integer | Truncate prompt to N tokens if it exceeds limit |
|
||||
|
||||
```python
|
||||
response = litellm.completion(
|
||||
model="sagemaker_nova/my-nova-endpoint",
|
||||
messages=[{"role": "user", "content": "Think step by step: what is 2+2?"}],
|
||||
top_k=40,
|
||||
reasoning_effort="low",
|
||||
logprobs=True,
|
||||
top_logprobs=2,
|
||||
)
|
||||
```
|
||||
|
||||
+3
-6
@@ -1464,12 +1464,9 @@ if TYPE_CHECKING:
|
||||
from .llms.petals.completion.transformation import PetalsConfig as PetalsConfig
|
||||
from .llms.ollama.chat.transformation import OllamaChatConfig as OllamaChatConfig
|
||||
from .llms.ollama.completion.transformation import OllamaConfig as OllamaConfig
|
||||
from .llms.sagemaker.completion.transformation import (
|
||||
SagemakerConfig as SagemakerConfig,
|
||||
)
|
||||
from .llms.sagemaker.chat.transformation import (
|
||||
SagemakerChatConfig as SagemakerChatConfig,
|
||||
)
|
||||
from .llms.sagemaker.completion.transformation import SagemakerConfig as SagemakerConfig
|
||||
from .llms.sagemaker.chat.transformation import SagemakerChatConfig as SagemakerChatConfig
|
||||
from .llms.sagemaker.nova.transformation import SagemakerNovaConfig as SagemakerNovaConfig
|
||||
from .llms.cohere.chat.transformation import CohereChatConfig as CohereChatConfig
|
||||
from .llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig as AnthropicMessagesConfig,
|
||||
|
||||
@@ -167,6 +167,7 @@ LLM_CONFIG_NAMES = (
|
||||
"OllamaConfig",
|
||||
"SagemakerConfig",
|
||||
"SagemakerChatConfig",
|
||||
"SagemakerNovaConfig",
|
||||
"CohereChatConfig",
|
||||
"AnthropicMessagesConfig",
|
||||
"AmazonAnthropicClaudeMessagesConfig",
|
||||
@@ -701,6 +702,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
|
||||
".llms.sagemaker.chat.transformation",
|
||||
"SagemakerChatConfig",
|
||||
),
|
||||
"SagemakerNovaConfig": (
|
||||
".llms.sagemaker.nova.transformation",
|
||||
"SagemakerNovaConfig",
|
||||
),
|
||||
"CohereChatConfig": (".llms.cohere.chat.transformation", "CohereChatConfig"),
|
||||
"AnthropicMessagesConfig": (
|
||||
".llms.anthropic.experimental_pass_through.messages.transformation",
|
||||
|
||||
@@ -505,6 +505,7 @@ LITELLM_CHAT_PROVIDERS = [
|
||||
"azure_ai",
|
||||
"sagemaker",
|
||||
"sagemaker_chat",
|
||||
"sagemaker_nova",
|
||||
"bedrock",
|
||||
"vllm",
|
||||
"nlp_cloud",
|
||||
|
||||
@@ -160,7 +160,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM):
|
||||
streaming_response = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="sagemaker_chat",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
return streaming_response
|
||||
@@ -180,8 +180,12 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM):
|
||||
signed_json_body: Optional[bytes] = None,
|
||||
) -> CustomStreamWrapper:
|
||||
if client is None or isinstance(client, HTTPHandler):
|
||||
try:
|
||||
llm_provider = LlmProviders(custom_llm_provider)
|
||||
except ValueError:
|
||||
llm_provider = LlmProviders.SAGEMAKER_CHAT
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=LlmProviders.SAGEMAKER_CHAT, params={}
|
||||
llm_provider=llm_provider, params={}
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -210,7 +214,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM):
|
||||
streaming_response = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="sagemaker_chat",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
return streaming_response
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
from .transformation import SagemakerNovaConfig # noqa: F401
|
||||
@@ -0,0 +1,70 @@
|
||||
"""
|
||||
Translate from OpenAI's `/v1/chat/completions` to SageMaker Nova Inference endpoints.
|
||||
|
||||
Nova models on SageMaker use OpenAI-compatible request/response format with
|
||||
additional Nova-specific parameters (top_k, reasoning_effort, etc.).
|
||||
|
||||
Docs: https://docs.aws.amazon.com/nova/latest/nova2-userguide/nova-sagemaker-inference-api-reference.html
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
from ..chat.transformation import SagemakerChatConfig
|
||||
|
||||
|
||||
class SagemakerNovaConfig(SagemakerChatConfig):
|
||||
"""
|
||||
Config for Amazon Nova models deployed on SageMaker Inference endpoints.
|
||||
|
||||
Nova uses OpenAI-compatible format (same as sagemaker_chat / HF Messages API)
|
||||
but with additional Nova-specific parameters and requires `stream: true` in
|
||||
the request body for streaming.
|
||||
|
||||
Usage:
|
||||
model="sagemaker_nova/<endpoint-name>"
|
||||
"""
|
||||
|
||||
@property
|
||||
def supports_stream_param_in_request_body(self) -> bool:
|
||||
"""Nova expects `stream: true` in the request body for streaming."""
|
||||
return True
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List:
|
||||
"""Extend parent params with Nova-specific parameters."""
|
||||
params = super().get_supported_openai_params(model)
|
||||
nova_params = [
|
||||
"top_k",
|
||||
"reasoning_effort",
|
||||
"allowed_token_ids",
|
||||
"truncate_prompt_tokens",
|
||||
]
|
||||
for p in nova_params:
|
||||
if p not in params:
|
||||
params.append(p)
|
||||
return params
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Nova SageMaker endpoints do not accept 'model' in the request body.
|
||||
Only supported fields: messages, max_tokens, max_completion_tokens,
|
||||
temperature, top_p, top_k, stream, stream_options, logprobs,
|
||||
top_logprobs, reasoning_effort, allowed_token_ids, truncate_prompt_tokens.
|
||||
"""
|
||||
request_body = super().transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
request_body.pop("model", None)
|
||||
return request_body
|
||||
+4
-2
@@ -3712,8 +3712,10 @@ def completion( # type: ignore # noqa: PLR0915
|
||||
):
|
||||
return _model_response
|
||||
response = _model_response
|
||||
elif custom_llm_provider == "sagemaker_chat":
|
||||
elif custom_llm_provider in ("sagemaker_chat", "sagemaker_nova"):
|
||||
# boto3 reads keys from .env
|
||||
# sagemaker_chat: HF Messages API endpoints
|
||||
# sagemaker_nova: Nova models on SageMaker (OpenAI-compatible)
|
||||
model_response = base_llm_http_handler.completion(
|
||||
model=model,
|
||||
stream=stream,
|
||||
@@ -3723,7 +3725,7 @@ def completion( # type: ignore # noqa: PLR0915
|
||||
model_response=model_response,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
custom_llm_provider="sagemaker_chat",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
timeout=timeout,
|
||||
headers=headers,
|
||||
encoding=_get_encoding(),
|
||||
|
||||
@@ -3131,6 +3131,7 @@ class LlmProviders(str, Enum):
|
||||
AZURE_AI = "azure_ai"
|
||||
SAGEMAKER = "sagemaker"
|
||||
SAGEMAKER_CHAT = "sagemaker_chat"
|
||||
SAGEMAKER_NOVA = "sagemaker_nova"
|
||||
BEDROCK = "bedrock"
|
||||
VLLM = "vllm"
|
||||
NLP_CLOUD = "nlp_cloud"
|
||||
|
||||
@@ -7945,6 +7945,7 @@ class ProviderConfigManager:
|
||||
LlmProviders.VERTEX_AI_BETA: (lambda: litellm.VertexGeminiConfig(), False),
|
||||
LlmProviders.CLOUDFLARE: (lambda: litellm.CloudflareChatConfig(), False),
|
||||
LlmProviders.SAGEMAKER_CHAT: (lambda: litellm.SagemakerChatConfig(), False),
|
||||
LlmProviders.SAGEMAKER_NOVA: (lambda: litellm.SagemakerNovaConfig(), False),
|
||||
LlmProviders.SAGEMAKER: (lambda: litellm.SagemakerConfig(), False),
|
||||
LlmProviders.FIREWORKS_AI: (lambda: litellm.FireworksAIConfig(), False),
|
||||
LlmProviders.FRIENDLIAI: (lambda: litellm.FriendliaiChatConfig(), False),
|
||||
|
||||
@@ -0,0 +1,276 @@
|
||||
"""
|
||||
Integration tests for SageMaker Nova provider.
|
||||
|
||||
These tests require a live SageMaker Nova endpoint and AWS credentials.
|
||||
They are skipped by default — run manually with:
|
||||
|
||||
pytest tests/test_litellm/llms/sagemaker/test_sagemaker_nova_integration.py -v --no-header -rN
|
||||
|
||||
Prerequisites:
|
||||
export AWS_PROFILE=<your-profile> # or set AWS_ACCESS_KEY_ID / AWS_SECRET_ACCESS_KEY
|
||||
export AWS_REGION_NAME=us-east-1
|
||||
export SAGEMAKER_NOVA_ENDPOINT=<your-endpoint-name>
|
||||
"""
|
||||
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import struct
|
||||
import zlib
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
ENDPOINT = os.environ.get("SAGEMAKER_NOVA_ENDPOINT", "")
|
||||
MODEL = f"sagemaker_nova/{ENDPOINT}"
|
||||
|
||||
skip_if_no_endpoint = pytest.mark.skipif(
|
||||
not ENDPOINT,
|
||||
reason="SAGEMAKER_NOVA_ENDPOINT not set — skipping live integration tests",
|
||||
)
|
||||
|
||||
|
||||
def _make_test_png() -> str:
|
||||
"""Create a minimal 4x4 PNG (red border, blue center) and return base64."""
|
||||
|
||||
def chunk(ctype, data):
|
||||
c = ctype + data
|
||||
return (
|
||||
struct.pack(">I", len(data))
|
||||
+ c
|
||||
+ struct.pack(">I", zlib.crc32(c) & 0xFFFFFFFF)
|
||||
)
|
||||
|
||||
width, height = 4, 4
|
||||
pixels = []
|
||||
for y in range(height):
|
||||
for x in range(width):
|
||||
if 1 <= x <= 2 and 1 <= y <= 2:
|
||||
pixels.append((0, 0, 255))
|
||||
else:
|
||||
pixels.append((255, 0, 0))
|
||||
|
||||
raw = b""
|
||||
for y in range(height):
|
||||
raw += b"\x00"
|
||||
for x in range(width):
|
||||
raw += bytes(pixels[y * width + x])
|
||||
|
||||
png = (
|
||||
b"\x89PNG\r\n\x1a\n"
|
||||
+ chunk(
|
||||
b"IHDR", struct.pack(">IIBBBBB", width, height, 8, 2, 0, 0, 0)
|
||||
)
|
||||
+ chunk(b"IDAT", zlib.compress(raw))
|
||||
+ chunk(b"IEND", b"")
|
||||
)
|
||||
return base64.b64encode(png).decode()
|
||||
|
||||
|
||||
@skip_if_no_endpoint
|
||||
class TestSagemakerNovaIntegration:
|
||||
"""Live integration tests for sagemaker_nova provider."""
|
||||
|
||||
def test_should_complete_basic_single_turn(self):
|
||||
"""Basic single-turn chat completion."""
|
||||
response = litellm.completion(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "What is 2+2? Reply in one word."}],
|
||||
max_tokens=32,
|
||||
temperature=0.1,
|
||||
)
|
||||
assert response.choices[0].message.content is not None
|
||||
assert len(response.choices[0].message.content.strip()) > 0
|
||||
assert response.choices[0].finish_reason == "stop"
|
||||
assert response.usage.prompt_tokens > 0
|
||||
assert response.usage.completion_tokens > 0
|
||||
assert response.usage.total_tokens == (
|
||||
response.usage.prompt_tokens + response.usage.completion_tokens
|
||||
)
|
||||
|
||||
def test_should_complete_multi_turn_conversation(self):
|
||||
"""Multi-turn conversation maintains context."""
|
||||
messages = [
|
||||
{"role": "user", "content": "My name is Alice."},
|
||||
]
|
||||
response1 = litellm.completion(
|
||||
model=MODEL,
|
||||
messages=messages,
|
||||
max_tokens=64,
|
||||
temperature=0.1,
|
||||
)
|
||||
assistant_msg = response1.choices[0].message.content
|
||||
assert assistant_msg is not None
|
||||
|
||||
# Second turn — model should remember the name
|
||||
messages.append({"role": "assistant", "content": assistant_msg})
|
||||
messages.append({"role": "user", "content": "What is my name?"})
|
||||
|
||||
response2 = litellm.completion(
|
||||
model=MODEL,
|
||||
messages=messages,
|
||||
max_tokens=64,
|
||||
temperature=0.1,
|
||||
)
|
||||
answer = response2.choices[0].message.content.lower()
|
||||
assert "alice" in answer, f"Expected 'alice' in response, got: {answer}"
|
||||
|
||||
def test_should_stream_response(self):
|
||||
"""Streaming returns chunks with content and final usage."""
|
||||
response = litellm.completion(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "Count from 1 to 5."}],
|
||||
max_tokens=64,
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
)
|
||||
|
||||
chunks = []
|
||||
full_content = ""
|
||||
for chunk in response:
|
||||
chunks.append(chunk)
|
||||
delta = chunk.choices[0].delta.content or ""
|
||||
full_content += delta
|
||||
|
||||
assert len(chunks) > 1, "Expected multiple streaming chunks"
|
||||
assert len(full_content.strip()) > 0, "Expected non-empty streamed content"
|
||||
|
||||
# Last chunk should have finish_reason
|
||||
final_chunks_with_finish = [
|
||||
c for c in chunks if c.choices and c.choices[0].finish_reason is not None
|
||||
]
|
||||
assert len(final_chunks_with_finish) > 0, "Expected at least one chunk with finish_reason"
|
||||
|
||||
def test_should_return_logprobs(self):
|
||||
"""Logprobs are returned when requested."""
|
||||
response = litellm.completion(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "Say hello."}],
|
||||
max_tokens=16,
|
||||
temperature=0.1,
|
||||
logprobs=True,
|
||||
top_logprobs=3,
|
||||
)
|
||||
lp = response.choices[0].logprobs
|
||||
assert lp is not None, "Expected logprobs in response"
|
||||
|
||||
content = lp.content if hasattr(lp, "content") else lp.get("content")
|
||||
assert content is not None and len(content) > 0, "Expected logprobs content"
|
||||
|
||||
first_token = content[0]
|
||||
assert "token" in first_token or hasattr(first_token, "token")
|
||||
assert "logprob" in first_token or hasattr(first_token, "logprob")
|
||||
|
||||
top = first_token.get("top_logprobs") if isinstance(first_token, dict) else first_token.top_logprobs
|
||||
assert top is not None and len(top) == 3, "Expected 3 top_logprobs"
|
||||
|
||||
def test_should_handle_multimodal_image_input(self):
|
||||
"""Multimodal with base64 image in content array."""
|
||||
b64_image = _make_test_png()
|
||||
response = litellm.completion(
|
||||
model=MODEL,
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "What colors do you see in this image? List them.",
|
||||
},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": f"data:image/png;base64,{b64_image}"
|
||||
},
|
||||
},
|
||||
],
|
||||
}
|
||||
],
|
||||
max_tokens=128,
|
||||
)
|
||||
content = response.choices[0].message.content.lower()
|
||||
assert response.choices[0].message.content is not None
|
||||
assert len(content) > 0
|
||||
# The image has red and blue — model should mention at least one
|
||||
assert "red" in content or "blue" in content, (
|
||||
f"Expected 'red' or 'blue' in multimodal response, got: {content}"
|
||||
)
|
||||
|
||||
def test_should_pass_nova_specific_params(self):
|
||||
"""Nova-specific parameters (top_k) are accepted."""
|
||||
response = litellm.completion(
|
||||
model=MODEL,
|
||||
messages=[{"role": "user", "content": "Say hello."}],
|
||||
max_tokens=32,
|
||||
top_k=40,
|
||||
temperature=0.7,
|
||||
)
|
||||
assert response.choices[0].message.content is not None
|
||||
assert response.usage.total_tokens > 0
|
||||
|
||||
def test_should_respect_system_message(self):
|
||||
"""System message should influence the response."""
|
||||
response = litellm.completion(
|
||||
model=MODEL,
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are a pirate. Always respond in pirate speak.",
|
||||
},
|
||||
{"role": "user", "content": "How are you today?"},
|
||||
],
|
||||
max_tokens=128,
|
||||
temperature=0.7,
|
||||
)
|
||||
content = response.choices[0].message.content.lower()
|
||||
assert response.choices[0].message.content is not None
|
||||
# Pirate-themed words likely in response
|
||||
pirate_words = ["arr", "ahoy", "matey", "ye", "sail", "sea", "cap"]
|
||||
assert any(
|
||||
w in content for w in pirate_words
|
||||
), f"Expected pirate speak, got: {content}"
|
||||
|
||||
|
||||
NOVA2_ENDPOINT = os.environ.get("SAGEMAKER_NOVA2_LITE_ENDPOINT", "")
|
||||
NOVA2_MODEL = f"sagemaker_nova/{NOVA2_ENDPOINT}"
|
||||
|
||||
skip_if_no_nova2_endpoint = pytest.mark.skipif(
|
||||
not NOVA2_ENDPOINT,
|
||||
reason="SAGEMAKER_NOVA2_LITE_ENDPOINT not set — requires Nova 2 Lite endpoint",
|
||||
)
|
||||
|
||||
|
||||
@skip_if_no_nova2_endpoint
|
||||
class TestSagemakerNova2LiteIntegration:
|
||||
"""
|
||||
Integration tests requiring a Nova 2 Lite endpoint (reasoning_effort support).
|
||||
|
||||
Run with:
|
||||
export SAGEMAKER_NOVA2_LITE_ENDPOINT=<your-nova-2-lite-endpoint>
|
||||
pytest tests/test_litellm/llms/sagemaker/test_sagemaker_nova_integration.py::TestSagemakerNova2LiteIntegration -v
|
||||
"""
|
||||
|
||||
def test_should_accept_reasoning_effort_low(self):
|
||||
"""reasoning_effort='low' should be accepted by Nova 2 Lite."""
|
||||
response = litellm.completion(
|
||||
model=NOVA2_MODEL,
|
||||
messages=[{"role": "user", "content": "What is 2+2?"}],
|
||||
max_tokens=32,
|
||||
reasoning_effort="low",
|
||||
)
|
||||
assert response.choices[0].message.content is not None
|
||||
assert response.usage.total_tokens > 0
|
||||
|
||||
def test_should_accept_reasoning_effort_high(self):
|
||||
"""reasoning_effort='high' should be accepted by Nova 2 Lite."""
|
||||
response = litellm.completion(
|
||||
model=NOVA2_MODEL,
|
||||
messages=[{"role": "user", "content": "Explain why the sky is blue."}],
|
||||
max_tokens=256,
|
||||
reasoning_effort="high",
|
||||
)
|
||||
assert response.choices[0].message.content is not None
|
||||
assert len(response.choices[0].message.content) > 0
|
||||
assert response.usage.completion_tokens > 0
|
||||
@@ -0,0 +1,393 @@
|
||||
"""
|
||||
Unit tests for SageMaker Nova transformation config.
|
||||
"""
|
||||
|
||||
import json
|
||||
import pytest
|
||||
|
||||
from litellm.llms.sagemaker.nova.transformation import SagemakerNovaConfig
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import convert_to_model_response_object
|
||||
|
||||
|
||||
class TestSagemakerNovaConfig:
|
||||
def setup_method(self):
|
||||
self.config = SagemakerNovaConfig()
|
||||
|
||||
def test_should_support_stream_param_in_request_body(self):
|
||||
"""Nova requires stream: true in the request body."""
|
||||
assert self.config.supports_stream_param_in_request_body is True
|
||||
|
||||
def test_should_include_nova_specific_params(self):
|
||||
"""Nova-specific params should be in the supported params list."""
|
||||
params = self.config.get_supported_openai_params(model="my-nova-endpoint")
|
||||
assert "top_k" in params
|
||||
assert "reasoning_effort" in params
|
||||
assert "allowed_token_ids" in params
|
||||
assert "truncate_prompt_tokens" in params
|
||||
|
||||
def test_should_include_standard_openai_params(self):
|
||||
"""Standard OpenAI params from parent should still be present."""
|
||||
params = self.config.get_supported_openai_params(model="my-nova-endpoint")
|
||||
assert "temperature" in params
|
||||
assert "max_tokens" in params
|
||||
assert "top_p" in params
|
||||
assert "stream" in params
|
||||
assert "logprobs" in params
|
||||
assert "top_logprobs" in params
|
||||
assert "stream_options" in params
|
||||
|
||||
def test_should_map_nova_params_to_request(self):
|
||||
"""Nova-specific params should pass through to optional_params."""
|
||||
optional_params = self.config.map_openai_params(
|
||||
non_default_params={
|
||||
"top_k": 40,
|
||||
"reasoning_effort": "low",
|
||||
"temperature": 0.7,
|
||||
},
|
||||
optional_params={},
|
||||
model="my-nova-endpoint",
|
||||
drop_params=False,
|
||||
)
|
||||
assert optional_params["top_k"] == 40
|
||||
assert optional_params["reasoning_effort"] == "low"
|
||||
assert optional_params["temperature"] == 0.7
|
||||
|
||||
def test_should_generate_correct_url_non_streaming(self):
|
||||
"""Non-streaming URL should use /invocations."""
|
||||
url = self.config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="my-nova-endpoint",
|
||||
optional_params={"aws_region_name": "us-east-1"},
|
||||
litellm_params={},
|
||||
stream=False,
|
||||
)
|
||||
assert url == "https://runtime.sagemaker.us-east-1.amazonaws.com/endpoints/my-nova-endpoint/invocations"
|
||||
|
||||
def test_should_generate_correct_url_streaming(self):
|
||||
"""Streaming URL should use /invocations-response-stream."""
|
||||
url = self.config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="my-nova-endpoint",
|
||||
optional_params={"aws_region_name": "us-east-1"},
|
||||
litellm_params={},
|
||||
stream=True,
|
||||
)
|
||||
assert url == "https://runtime.sagemaker.us-east-1.amazonaws.com/endpoints/my-nova-endpoint/invocations-response-stream"
|
||||
|
||||
def test_should_have_custom_stream_wrapper(self):
|
||||
"""Nova should use custom stream wrapper (AWS EventStream)."""
|
||||
assert self.config.has_custom_stream_wrapper is True
|
||||
|
||||
|
||||
class TestSagemakerNovaResponseParsing:
|
||||
"""Test that Nova's OpenAI-compatible responses are correctly parsed."""
|
||||
|
||||
def test_should_parse_non_streaming_response(self):
|
||||
"""Nova non-streaming response should be parsed into ModelResponse."""
|
||||
nova_response = {
|
||||
"id": "chatcmpl-123e4567-e89b-12d3-a456-426614174000",
|
||||
"object": "chat.completion",
|
||||
"created": 1677652288,
|
||||
"model": "nova-micro-custom",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Hello! How can I help?",
|
||||
"refusal": None,
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 9,
|
||||
"completion_tokens": 12,
|
||||
"total_tokens": 21,
|
||||
},
|
||||
}
|
||||
result = convert_to_model_response_object(
|
||||
response_object=nova_response,
|
||||
model_response_object=ModelResponse(),
|
||||
)
|
||||
assert result.id == "chatcmpl-123e4567-e89b-12d3-a456-426614174000"
|
||||
assert result.choices[0].message.content == "Hello! How can I help?"
|
||||
assert result.choices[0].finish_reason == "stop"
|
||||
assert result.usage.prompt_tokens == 9
|
||||
assert result.usage.completion_tokens == 12
|
||||
assert result.usage.total_tokens == 21
|
||||
|
||||
def test_should_parse_response_with_reasoning_content(self):
|
||||
"""Nova reasoning_content should be extracted correctly."""
|
||||
nova_response = {
|
||||
"id": "chatcmpl-reasoning-test",
|
||||
"object": "chat.completion",
|
||||
"created": 1677652288,
|
||||
"model": "nova-2-lite-custom",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "The answer is 4.",
|
||||
"reasoning_content": "Let me think: 2+2=4",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 15,
|
||||
"completion_tokens": 20,
|
||||
"total_tokens": 35,
|
||||
},
|
||||
}
|
||||
result = convert_to_model_response_object(
|
||||
response_object=nova_response,
|
||||
model_response_object=ModelResponse(),
|
||||
)
|
||||
assert result.choices[0].message.content == "The answer is 4."
|
||||
assert result.choices[0].message.reasoning_content == "Let me think: 2+2=4"
|
||||
|
||||
def test_should_parse_response_with_logprobs(self):
|
||||
"""Nova logprobs should be preserved in response."""
|
||||
nova_response = {
|
||||
"id": "chatcmpl-logprobs-test",
|
||||
"object": "chat.completion",
|
||||
"created": 1677652288,
|
||||
"model": "nova-micro-custom",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Hello",
|
||||
},
|
||||
"logprobs": {
|
||||
"content": [
|
||||
{
|
||||
"token": "Hello",
|
||||
"logprob": -0.5,
|
||||
"top_logprobs": [
|
||||
{"token": "Hello", "logprob": -0.5},
|
||||
{"token": "Hi", "logprob": -1.2},
|
||||
],
|
||||
}
|
||||
]
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 5,
|
||||
"completion_tokens": 1,
|
||||
"total_tokens": 6,
|
||||
},
|
||||
}
|
||||
result = convert_to_model_response_object(
|
||||
response_object=nova_response,
|
||||
model_response_object=ModelResponse(),
|
||||
)
|
||||
assert result.choices[0].logprobs is not None
|
||||
assert result.choices[0].logprobs["content"][0]["token"] == "Hello"
|
||||
|
||||
def test_should_parse_response_with_cached_tokens(self):
|
||||
"""Nova prompt_tokens_details with cached_tokens should be parsed."""
|
||||
nova_response = {
|
||||
"id": "chatcmpl-cached-test",
|
||||
"object": "chat.completion",
|
||||
"created": 1677652288,
|
||||
"model": "nova-micro-custom",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Hi",
|
||||
},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 20,
|
||||
"completion_tokens": 1,
|
||||
"total_tokens": 21,
|
||||
"prompt_tokens_details": {"cached_tokens": 10},
|
||||
},
|
||||
}
|
||||
result = convert_to_model_response_object(
|
||||
response_object=nova_response,
|
||||
model_response_object=ModelResponse(),
|
||||
)
|
||||
assert result.usage.prompt_tokens_details.cached_tokens == 10
|
||||
|
||||
|
||||
class TestSagemakerChatBackwardsCompatibility:
|
||||
"""Verify that changes to SagemakerChatConfig don't break existing sagemaker_chat callers."""
|
||||
|
||||
def setup_method(self):
|
||||
from litellm.llms.sagemaker.chat.transformation import SagemakerChatConfig
|
||||
self.config = SagemakerChatConfig()
|
||||
|
||||
def test_should_not_support_stream_param_in_request_body(self):
|
||||
"""sagemaker_chat should NOT send stream in request body (unchanged behavior)."""
|
||||
assert self.config.supports_stream_param_in_request_body is False
|
||||
|
||||
def test_should_generate_correct_urls(self):
|
||||
"""sagemaker_chat URLs should be unchanged."""
|
||||
url = self.config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="my-hf-endpoint",
|
||||
optional_params={"aws_region_name": "us-west-2"},
|
||||
litellm_params={},
|
||||
stream=False,
|
||||
)
|
||||
assert url == "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/my-hf-endpoint/invocations"
|
||||
|
||||
stream_url = self.config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="my-hf-endpoint",
|
||||
optional_params={"aws_region_name": "us-west-2"},
|
||||
litellm_params={},
|
||||
stream=True,
|
||||
)
|
||||
assert stream_url == "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/my-hf-endpoint/invocations-response-stream"
|
||||
|
||||
def test_should_still_have_custom_stream_wrapper(self):
|
||||
"""sagemaker_chat should still use custom stream wrapper."""
|
||||
assert self.config.has_custom_stream_wrapper is True
|
||||
|
||||
def test_should_not_include_nova_specific_params(self):
|
||||
"""sagemaker_chat should NOT have Nova-specific params."""
|
||||
params = self.config.get_supported_openai_params(model="my-hf-endpoint")
|
||||
assert "top_k" not in params
|
||||
assert "reasoning_effort" not in params
|
||||
assert "allowed_token_ids" not in params
|
||||
assert "truncate_prompt_tokens" not in params
|
||||
|
||||
def test_should_preserve_standard_openai_params(self):
|
||||
"""sagemaker_chat should still support standard OpenAI params."""
|
||||
params = self.config.get_supported_openai_params(model="my-hf-endpoint")
|
||||
assert "temperature" in params
|
||||
assert "max_tokens" in params
|
||||
assert "top_p" in params
|
||||
assert "stream" in params
|
||||
|
||||
def test_sync_stream_wrapper_uses_correct_provider_string(self):
|
||||
"""
|
||||
Verify that when get_sync_custom_stream_wrapper is called with
|
||||
custom_llm_provider="sagemaker_chat", the CustomStreamWrapper
|
||||
receives "sagemaker_chat" (not something else).
|
||||
"""
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.iter_bytes.return_value = iter([])
|
||||
mock_client.post.return_value = mock_response
|
||||
|
||||
with patch("litellm.llms.sagemaker.chat.transformation.CustomStreamWrapper") as mock_csw:
|
||||
mock_csw.return_value = MagicMock()
|
||||
self.config.get_sync_custom_stream_wrapper(
|
||||
model="my-hf-endpoint",
|
||||
custom_llm_provider="sagemaker_chat",
|
||||
logging_obj=MagicMock(),
|
||||
api_base="https://example.com",
|
||||
headers={},
|
||||
data={},
|
||||
messages=[],
|
||||
client=mock_client,
|
||||
)
|
||||
mock_csw.assert_called_once()
|
||||
call_kwargs = mock_csw.call_args[1]
|
||||
assert call_kwargs["custom_llm_provider"] == "sagemaker_chat"
|
||||
|
||||
def test_async_stream_wrapper_uses_correct_provider_string(self):
|
||||
"""
|
||||
Verify that when get_async_custom_stream_wrapper is called with
|
||||
custom_llm_provider="sagemaker_chat", the CustomStreamWrapper
|
||||
receives "sagemaker_chat".
|
||||
"""
|
||||
import asyncio
|
||||
from unittest.mock import patch, MagicMock, AsyncMock
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
|
||||
async def empty_aiter():
|
||||
return
|
||||
yield # make it an async generator
|
||||
|
||||
mock_response.aiter_bytes.return_value = empty_aiter()
|
||||
mock_client.post.return_value = mock_response
|
||||
|
||||
with patch("litellm.llms.sagemaker.chat.transformation.CustomStreamWrapper") as mock_csw:
|
||||
mock_csw.return_value = MagicMock()
|
||||
asyncio.run(
|
||||
self.config.get_async_custom_stream_wrapper(
|
||||
model="my-hf-endpoint",
|
||||
custom_llm_provider="sagemaker_chat",
|
||||
logging_obj=MagicMock(),
|
||||
api_base="https://example.com",
|
||||
headers={},
|
||||
data={},
|
||||
messages=[],
|
||||
client=mock_client,
|
||||
)
|
||||
)
|
||||
mock_csw.assert_called_once()
|
||||
call_kwargs = mock_csw.call_args[1]
|
||||
assert call_kwargs["custom_llm_provider"] == "sagemaker_chat"
|
||||
|
||||
def test_async_stream_wrapper_llm_provider_enum_resolves(self):
|
||||
"""
|
||||
Verify LlmProviders(custom_llm_provider) resolves correctly for
|
||||
"sagemaker_chat" and doesn't fall through to the ValueError fallback.
|
||||
"""
|
||||
from litellm.types.utils import LlmProviders
|
||||
provider = LlmProviders("sagemaker_chat")
|
||||
assert provider == LlmProviders.SAGEMAKER_CHAT
|
||||
|
||||
|
||||
class TestSagemakerNovaTransformRequest:
|
||||
"""Test Nova-specific request transformation."""
|
||||
|
||||
def setup_method(self):
|
||||
self.config = SagemakerNovaConfig()
|
||||
|
||||
def test_should_not_include_model_in_request_body(self):
|
||||
"""Nova SageMaker endpoints reject 'model' in the request body."""
|
||||
request = self.config.transform_request(
|
||||
model="my-nova-endpoint",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={"temperature": 0.7},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert "model" not in request
|
||||
assert "messages" in request
|
||||
assert request["temperature"] == 0.7
|
||||
|
||||
def test_should_include_all_nova_params_in_request(self):
|
||||
"""Nova-specific params should appear in the request body."""
|
||||
request = self.config.transform_request(
|
||||
model="my-nova-endpoint",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
optional_params={
|
||||
"top_k": 40,
|
||||
"max_tokens": 512,
|
||||
"reasoning_effort": "low",
|
||||
},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert "model" not in request
|
||||
assert request["top_k"] == 40
|
||||
assert request["max_tokens"] == 512
|
||||
assert request["reasoning_effort"] == "low"
|
||||
Reference in New Issue
Block a user