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:
milan-berri 2026-05-22 22:00:42 +03:00 • committed by GitHub
parent 643989989f
commit 9600fda2cc
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 365 additions and 19 deletions

View file

@ -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,

View file

@ -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

View 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"}

View file

@ -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,

View file

@ -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

View file

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