mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(sagemaker): send native Cohere embed payload to Cohere SageMaker endpoints (#28613)
* fix(sagemaker): use Cohere embed payload for Marketplace endpoints
SageMaker embedding only special-cased Voyage; every other endpoint received
HuggingFace TGI `{"inputs": [...]}`. AWS Marketplace Cohere containers expect
the native Cohere embed payload (`texts`, `input_type`) and reject the HF
shape with `422 EmbedReqV2.inputs is of type string but should be of type
Object`.
Add `SagemakerCohereEmbeddingConfig` that reuses Bedrock/Cohere request and
response transforms, and route SageMaker endpoint names containing `cohere`
or a Cohere embed model fragment (`embed-multilingual`, `embed-english`,
`embed-v3`, `embed-v4`) to it. Supports `input_type`, `dimensions`, and
`encoding_format`. Voyage and HuggingFace SageMaker endpoints are unchanged.
Co-authored-by: Cursor <cursoragent@cursor.com>
* refactor(sagemaker): simplify cohere detection and align with file conventions
- Detect Cohere SageMaker endpoints with a single `"cohere" in model.lower()`
check, mirroring the existing Voyage branch instead of a separate helper
function and marker constant.
- Drop instance caches of sub-configs; instantiate `BedrockCohereEmbeddingConfig`
/ `CohereEmbeddingConfig` per call to match the existing pattern in
`BedrockCohereEmbeddingConfig._transform_request`.
- Match `SagemakerEmbeddingConfig`'s signatures, defaults, and `Any` typing for
`logging_obj`; collapse the input-normalization helper inline.
- Inline `transform_embedding_response` input lookup; no behavior change.
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(sagemaker): restore provider-supported embedding params after map
Cohere input_type is advertised in get_supported_openai_params but was
filtered out of non_default_params by OPENAI_EMBEDDING_PARAMS before
map_openai_params ran. Merge supported params from passed_params after
map (same path Greptile flagged). Handle input_type explicitly in
SagemakerCohereEmbeddingConfig.map_openai_params and add an integration
test through get_optional_params_embeddings.
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(embeddings): only restore non-OpenAI supported params after map
The post-map restore loop must skip OPENAI_EMBEDDING_PARAMS so mapped
fields (e.g. dimensions -> output_dimension) are not duplicated under
their OpenAI names. Align SageMaker embedding import order with sibling
files and add a regression test for dimensions mapping.
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(sagemaker): avoid double post_call on Cohere embedding response
Greptile review on #28613 caught that `CohereEmbeddingConfig._transform_response`
calls `logging_obj.post_call` internally. The SageMaker embedding handler
already calls `post_call` once before invoking the transform, so the Cohere
SageMaker path fired callbacks, cost calculators, and log handlers twice
per request.
Extract the parsing body of `_transform_response` into
`_populate_embedding_response` (pure extract-method, no behavior change
for existing Cohere direct or Bedrock Cohere paths, which keep calling
`_transform_response`). Have `SagemakerCohereEmbeddingConfig` call the
new helper directly so it parses the response without re-logging.
Add a regression test asserting `logging_obj.post_call` is not invoked
by the SageMaker Cohere transform.
Co-authored-by: Cursor <cursoragent@cursor.com>
---------
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
643989989f
commit
9600fda2cc
6 changed files with 365 additions and 19 deletions
|
|
@ -110,15 +110,35 @@ class CohereEmbeddingConfig:
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=response_json,
|
||||
)
|
||||
return self._populate_embedding_response(
|
||||
response_json=response_json,
|
||||
model_response=model_response,
|
||||
model=model,
|
||||
encoding=encoding,
|
||||
input=input,
|
||||
)
|
||||
|
||||
def _populate_embedding_response(
|
||||
self,
|
||||
response_json: dict,
|
||||
model_response: EmbeddingResponse,
|
||||
model: str,
|
||||
encoding: Any,
|
||||
input: list,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
response
|
||||
Parse a Cohere embed response body into an OpenAI-style EmbeddingResponse.
|
||||
|
||||
Split out from `_transform_response` so callers that already log
|
||||
`post_call` themselves (e.g. SageMaker's embedding handler) can reuse
|
||||
the parsing without triggering a second `post_call`.
|
||||
|
||||
Response shape:
|
||||
{
|
||||
'object': "list",
|
||||
'data': [
|
||||
|
||||
]
|
||||
'model',
|
||||
'usage'
|
||||
'data': [...],
|
||||
'model',
|
||||
'usage',
|
||||
}
|
||||
"""
|
||||
embeddings = response_json["embeddings"]
|
||||
|
|
@ -149,9 +169,6 @@ class CohereEmbeddingConfig:
|
|||
model_response.object = "list"
|
||||
model_response.data = output_data
|
||||
model_response.model = model
|
||||
input_tokens = 0
|
||||
for text in input:
|
||||
input_tokens += len(encoding.encode(text))
|
||||
|
||||
setattr(
|
||||
model_response,
|
||||
|
|
|
|||
|
|
@ -578,7 +578,7 @@ class SagemakerLLM(BaseAWSLLM):
|
|||
logger_fn=None,
|
||||
):
|
||||
"""
|
||||
Supports both Huggingface Jumpstart embeddings and Voyage models
|
||||
Supports Hugging Face (TGI), Voyage, and Cohere embedding endpoints
|
||||
"""
|
||||
### BOTO3 INIT
|
||||
import boto3
|
||||
|
|
|
|||
141
litellm/llms/sagemaker/embedding/cohere_transformation.py
Normal file
141
litellm/llms/sagemaker/embedding/cohere_transformation.py
Normal file
|
|
@ -0,0 +1,141 @@
|
|||
"""
|
||||
Translate from OpenAI's `/v1/embeddings` to Sagemaker's `/invoke`
|
||||
|
||||
In the native Cohere embed format for self-hosted Cohere endpoints
|
||||
(AWS Marketplace / JumpStart). Cohere containers expect
|
||||
`{"texts": [...], "input_type": "..."}` and reject the HuggingFace TGI shape
|
||||
`{"inputs": [...]}` with `422 EmbedReqV2.inputs is of type string but should
|
||||
be of type Object`.
|
||||
|
||||
Reference: https://docs.cohere.com/v2/reference/embed
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Union, cast
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.openai import AllEmbeddingInputValues
|
||||
|
||||
from httpx._models import Headers, Response
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.llms.bedrock.embed.cohere_transformation import (
|
||||
BedrockCohereEmbeddingConfig,
|
||||
)
|
||||
from litellm.llms.cohere.embed.v1_transformation import CohereEmbeddingConfig
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
from ..common_utils import SagemakerError
|
||||
|
||||
|
||||
class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig):
|
||||
"""
|
||||
SageMaker invoke payload for self-hosted Cohere embed models.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
return ["encoding_format", "dimensions", "input_type"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
optional_params = BedrockCohereEmbeddingConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
if "input_type" in non_default_params:
|
||||
optional_params["input_type"] = non_default_params["input_type"]
|
||||
return optional_params
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, Headers]
|
||||
) -> BaseLLMException:
|
||||
return SagemakerError(
|
||||
message=error_message, status_code=status_code, headers=headers
|
||||
)
|
||||
|
||||
def transform_embedding_request(
|
||||
self,
|
||||
model: str,
|
||||
input: "AllEmbeddingInputValues",
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform embedding request for Cohere models on SageMaker
|
||||
"""
|
||||
if isinstance(input, str):
|
||||
input_list: List[str] = [input]
|
||||
elif isinstance(input, list):
|
||||
if input and (isinstance(input[0], list) or isinstance(input[0], int)):
|
||||
raise ValueError("Input must be a list of strings")
|
||||
input_list = cast(List[str], input)
|
||||
else:
|
||||
input_list = [str(input)]
|
||||
|
||||
return dict(
|
||||
BedrockCohereEmbeddingConfig()._transform_request(
|
||||
model=model,
|
||||
input=input_list,
|
||||
inference_params=optional_params,
|
||||
)
|
||||
)
|
||||
|
||||
def transform_embedding_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: Response,
|
||||
model_response: "EmbeddingResponse",
|
||||
logging_obj: Any,
|
||||
api_key: Optional[str] = None,
|
||||
request_data: dict = {},
|
||||
optional_params: dict = {},
|
||||
litellm_params: dict = {},
|
||||
) -> "EmbeddingResponse":
|
||||
"""
|
||||
Transform embedding response for Cohere models on SageMaker.
|
||||
|
||||
Uses `CohereEmbeddingConfig._populate_embedding_response` (not
|
||||
`_transform_response`) so we do not log `post_call` a second time
|
||||
— the SageMaker embedding handler already logs `post_call` before
|
||||
invoking this transform.
|
||||
"""
|
||||
input_value = (
|
||||
logging_obj.model_call_details.get("input")
|
||||
or request_data.get("texts")
|
||||
or request_data.get("images")
|
||||
or []
|
||||
)
|
||||
if isinstance(input_value, str):
|
||||
input_value = [input_value]
|
||||
|
||||
return CohereEmbeddingConfig()._populate_embedding_response(
|
||||
response_json=raw_response.json(),
|
||||
model_response=model_response,
|
||||
model=model,
|
||||
encoding=litellm.encoding,
|
||||
input=input_value,
|
||||
)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[Any],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Validate environment for SageMaker Cohere embeddings
|
||||
"""
|
||||
return {"Content-Type": "application/json"}
|
||||
|
|
@ -11,12 +11,13 @@ if TYPE_CHECKING:
|
|||
|
||||
from httpx._models import Headers, Response
|
||||
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.utils import Usage, EmbeddingResponse
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.llms.voyage.embedding.transformation import VoyageEmbeddingConfig
|
||||
from litellm.types.utils import EmbeddingResponse, Usage
|
||||
|
||||
from ..common_utils import SagemakerError
|
||||
from .cohere_transformation import SagemakerCohereEmbeddingConfig
|
||||
|
||||
|
||||
class SagemakerEmbeddingConfig(BaseEmbeddingConfig):
|
||||
|
|
@ -38,17 +39,20 @@ class SagemakerEmbeddingConfig(BaseEmbeddingConfig):
|
|||
Returns:
|
||||
Appropriate embedding config instance
|
||||
"""
|
||||
if "voyage" in model.lower():
|
||||
model_lower = model.lower()
|
||||
if "voyage" in model_lower:
|
||||
return VoyageEmbeddingConfig()
|
||||
else:
|
||||
return cls()
|
||||
if "cohere" in model_lower:
|
||||
return SagemakerCohereEmbeddingConfig()
|
||||
return cls()
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
# Check if this is an embedding model
|
||||
if "voyage" in model.lower():
|
||||
model_lower = model.lower()
|
||||
if "voyage" in model_lower:
|
||||
return VoyageEmbeddingConfig().get_supported_openai_params(model)
|
||||
else:
|
||||
return []
|
||||
if "cohere" in model_lower:
|
||||
return SagemakerCohereEmbeddingConfig().get_supported_openai_params(model)
|
||||
return []
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -3350,6 +3350,21 @@ def get_optional_params_embeddings( # noqa: PLR0915
|
|||
model=model,
|
||||
drop_params=drop_params if drop_params is not None else False,
|
||||
)
|
||||
# Provider-only params (e.g. Cohere input_type) are not in
|
||||
# OPENAI_EMBEDDING_PARAMS, so embedding_pre_process drops them from
|
||||
# non_default_params before map_openai_params. Restore only those extras
|
||||
# from passed_params — skip OPENAI_EMBEDDING_PARAMS to avoid duplicating
|
||||
# values already mapped (e.g. dimensions -> output_dimension).
|
||||
if supported_params:
|
||||
for param in supported_params:
|
||||
if param in OPENAI_EMBEDDING_PARAMS:
|
||||
continue
|
||||
if (
|
||||
param in passed_params
|
||||
and passed_params[param] is not None
|
||||
and param not in optional_params
|
||||
):
|
||||
optional_params[param] = passed_params[param]
|
||||
## raise exception if non-default value passed for non-openai/azure embedding calls
|
||||
elif custom_llm_provider == "openai":
|
||||
# 'dimensions` is only supported in `text-embedding-3` and later models
|
||||
|
|
|
|||
|
|
@ -17,6 +17,9 @@ import pytest
|
|||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
||||
from litellm import embedding
|
||||
from litellm.llms.sagemaker.embedding.cohere_transformation import (
|
||||
SagemakerCohereEmbeddingConfig,
|
||||
)
|
||||
from litellm.llms.sagemaker.embedding.transformation import SagemakerEmbeddingConfig
|
||||
from litellm.llms.voyage.embedding.transformation import VoyageEmbeddingConfig
|
||||
from litellm.types.utils import EmbeddingResponse, Usage
|
||||
|
|
@ -54,6 +57,172 @@ class TestSagemakerEmbeddingFactory:
|
|||
assert isinstance(config2, VoyageEmbeddingConfig)
|
||||
assert isinstance(config3, VoyageEmbeddingConfig)
|
||||
|
||||
def test_get_model_config_cohere_model(self):
|
||||
"""Cohere SageMaker endpoints route to SagemakerCohereEmbeddingConfig"""
|
||||
for endpoint_name in (
|
||||
"cohere.embed-multilingual-v3",
|
||||
"cohere-embed-english-v3-prod",
|
||||
"my-cohere-marketplace-endpoint",
|
||||
"COHERE-EMBED-V4",
|
||||
):
|
||||
config = SagemakerEmbeddingConfig.get_model_config(endpoint_name)
|
||||
assert isinstance(config, SagemakerCohereEmbeddingConfig), endpoint_name
|
||||
|
||||
|
||||
class TestSagemakerCohereEmbeddingConfig:
|
||||
"""Cohere-specific SageMaker embedding request/response transforms"""
|
||||
|
||||
def setup_method(self):
|
||||
self.config = SagemakerCohereEmbeddingConfig()
|
||||
|
||||
MODEL = "cohere.embed-multilingual-v3"
|
||||
|
||||
def test_transform_request_uses_cohere_payload(self):
|
||||
"""Bug repro: request must use `texts` + `input_type`, not HF `inputs`"""
|
||||
result = self.config.transform_embedding_request(
|
||||
model=self.MODEL,
|
||||
input=["hello"],
|
||||
optional_params={"input_type": "search_query"},
|
||||
headers={},
|
||||
)
|
||||
assert "inputs" not in result
|
||||
assert result["texts"] == ["hello"]
|
||||
assert result["input_type"] == "search_query"
|
||||
|
||||
def test_transform_request_default_input_type(self):
|
||||
result = self.config.transform_embedding_request(
|
||||
model=self.MODEL,
|
||||
input=["hello"],
|
||||
optional_params={},
|
||||
headers={},
|
||||
)
|
||||
assert result["texts"] == ["hello"]
|
||||
assert result["input_type"] == "search_document"
|
||||
|
||||
def test_transform_request_normalizes_string_input(self):
|
||||
result = self.config.transform_embedding_request(
|
||||
model=self.MODEL,
|
||||
input="hello",
|
||||
optional_params={},
|
||||
headers={},
|
||||
)
|
||||
assert result["texts"] == ["hello"]
|
||||
|
||||
def test_map_openai_params_dimensions_to_output_dimension(self):
|
||||
params = self.config.map_openai_params(
|
||||
non_default_params={"dimensions": 512, "encoding_format": "float"},
|
||||
optional_params={},
|
||||
model=self.MODEL,
|
||||
drop_params=False,
|
||||
)
|
||||
assert params["output_dimension"] == 512
|
||||
assert params["embedding_types"] == ["float"]
|
||||
|
||||
def test_map_openai_params_input_type_from_non_default_params(self):
|
||||
params = self.config.map_openai_params(
|
||||
non_default_params={"input_type": "search_query"},
|
||||
optional_params={},
|
||||
model=self.MODEL,
|
||||
drop_params=False,
|
||||
)
|
||||
assert params["input_type"] == "search_query"
|
||||
|
||||
def test_get_optional_params_embeddings_preserves_input_type(self):
|
||||
"""Exercises get_optional_params_embeddings, not transform in isolation."""
|
||||
from litellm.utils import get_optional_params_embeddings
|
||||
|
||||
optional_params = get_optional_params_embeddings(
|
||||
model=self.MODEL,
|
||||
custom_llm_provider="sagemaker",
|
||||
input_type="search_query",
|
||||
)
|
||||
assert optional_params.get("input_type") == "search_query"
|
||||
|
||||
body = self.config.transform_embedding_request(
|
||||
model=self.MODEL,
|
||||
input=["hello"],
|
||||
optional_params=optional_params,
|
||||
headers={},
|
||||
)
|
||||
assert body["texts"] == ["hello"]
|
||||
assert body["input_type"] == "search_query"
|
||||
|
||||
def test_get_optional_params_embeddings_maps_dimensions_without_duplicate(self):
|
||||
"""dimensions must map to output_dimension only, not also stay as dimensions."""
|
||||
from litellm.utils import get_optional_params_embeddings
|
||||
|
||||
optional_params = get_optional_params_embeddings(
|
||||
model=self.MODEL,
|
||||
custom_llm_provider="sagemaker",
|
||||
dimensions=512,
|
||||
input_type="search_query",
|
||||
)
|
||||
assert optional_params.get("output_dimension") == 512
|
||||
assert "dimensions" not in optional_params
|
||||
assert optional_params.get("input_type") == "search_query"
|
||||
|
||||
def test_transform_response_parses_cohere_payload(self):
|
||||
cohere_response = {
|
||||
"embeddings": [[0.1, 0.2, 0.3]],
|
||||
"meta": {"billed_units": {"input_tokens": 2}},
|
||||
}
|
||||
mock_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=json.dumps(cohere_response).encode("utf-8"),
|
||||
headers={"content-type": "application/json"},
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {"input": ["hello"]}
|
||||
|
||||
result = self.config.transform_embedding_response(
|
||||
model=self.MODEL,
|
||||
raw_response=mock_response,
|
||||
model_response=EmbeddingResponse(),
|
||||
logging_obj=logging_obj,
|
||||
api_key=None,
|
||||
request_data={"texts": ["hello"], "input_type": "search_query"},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert result.object == "list"
|
||||
assert len(result.data) == 1
|
||||
assert result.data[0]["embedding"] == [0.1, 0.2, 0.3]
|
||||
assert result.usage.prompt_tokens == 2
|
||||
|
||||
def test_transform_response_does_not_double_call_post_call(self):
|
||||
"""
|
||||
Greptile review fix: SageMaker handler already calls
|
||||
`logging_obj.post_call` once before invoking
|
||||
`transform_embedding_response`. The transform must NOT call it again,
|
||||
otherwise callbacks, cost calculators, and log handlers double-fire
|
||||
for every Cohere SageMaker embedding call.
|
||||
"""
|
||||
cohere_response = {
|
||||
"embeddings": [[0.1, 0.2, 0.3]],
|
||||
"meta": {"billed_units": {"input_tokens": 2}},
|
||||
}
|
||||
mock_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=json.dumps(cohere_response).encode("utf-8"),
|
||||
headers={"content-type": "application/json"},
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {"input": ["hello"]}
|
||||
|
||||
self.config.transform_embedding_response(
|
||||
model=self.MODEL,
|
||||
raw_response=mock_response,
|
||||
model_response=EmbeddingResponse(),
|
||||
logging_obj=logging_obj,
|
||||
api_key=None,
|
||||
request_data={"texts": ["hello"], "input_type": "search_query"},
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
logging_obj.post_call.assert_not_called()
|
||||
|
||||
|
||||
class TestVoyageEmbeddingConfig:
|
||||
"""Test Voyage-specific embedding configuration"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue