mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Add Volcengine embedding module with handler and transformation logic
- Implemented VolcEngineEmbeddingHandler for synchronous and asynchronous embedding requests. - Created VolcEngineEmbeddingConfig for transforming requests and responses to/from Volcengine format. - Added integration tests for embedding functionality, covering various scenarios including error handling and parameter validation. - Established test structure for Volcengine embedding, ensuring compliance with LiteLLM testing patterns. - Included comprehensive tests for parameter mapping, request transformation, and response handling.
This commit is contained in:
parent
6ce1d82970
commit
2d0a57a719
17 changed files with 1438 additions and 15 deletions
|
|
@ -3,7 +3,7 @@ https://www.volcengine.com/docs/82379/1263482
|
|||
|
||||
:::tip
|
||||
|
||||
**We support ALL Volcengine NIM models, just set `model=volcengine/<any-model-on-volcengine>` as a prefix when sending litellm requests**
|
||||
**We support ALL Volcengine models including Chat and Embeddings, just set `model=volcengine/<any-model-on-volcengine>` as a prefix when sending litellm requests**
|
||||
|
||||
:::
|
||||
|
||||
|
|
@ -11,6 +11,8 @@ https://www.volcengine.com/docs/82379/1263482
|
|||
```python
|
||||
# env variable
|
||||
os.environ['VOLCENGINE_API_KEY']
|
||||
# or
|
||||
os.environ['ARK_API_KEY']
|
||||
```
|
||||
|
||||
## Sample Usage
|
||||
|
|
@ -64,9 +66,42 @@ for chunk in response:
|
|||
print(chunk)
|
||||
```
|
||||
|
||||
## Sample Usage - Embedding
|
||||
```python
|
||||
from litellm import embedding
|
||||
import os
|
||||
|
||||
## Supported Models - 💥 ALL Volcengine NIM Models Supported!
|
||||
We support ALL `volcengine` models, just set `volcengine/<OUR_ENDPOINT_ID>` as a prefix when sending completion requests
|
||||
os.environ['VOLCENGINE_API_KEY'] = ""
|
||||
response = embedding(
|
||||
model="volcengine/doubao-embedding-text-240715",
|
||||
input=["hello world", "good morning"]
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
### Supported Embedding Models
|
||||
- `doubao-embedding-large` (2048 dimensions)
|
||||
- `doubao-embedding-large-text-250515` (2048 dimensions)
|
||||
- `doubao-embedding-large-text-240915` (4096 dimensions)
|
||||
- `doubao-embedding` (2560 dimensions)
|
||||
- `doubao-embedding-text-240715` (2560 dimensions)
|
||||
|
||||
### Embedding Parameters
|
||||
```python
|
||||
from litellm import embedding
|
||||
|
||||
response = embedding(
|
||||
model="volcengine/doubao-embedding-text-240715",
|
||||
input=["sample text"],
|
||||
encoding_format="float", # optional: "float" (default), "base64"
|
||||
user="user-123", # optional: user identifier for tracking
|
||||
)
|
||||
```
|
||||
|
||||
## Supported Models - 💥 ALL Volcengine Models Supported!
|
||||
We support ALL `volcengine` models for both chat completions and embeddings:
|
||||
- **Chat Models**: Set `volcengine/<OUR_ENDPOINT_ID>` as a prefix when sending completion requests
|
||||
- **Embedding Models**: Use the specific model names listed above (e.g., `volcengine/doubao-embedding-text-240715`)
|
||||
|
||||
## Sample Usage - LiteLLM Proxy
|
||||
|
||||
|
|
@ -74,14 +109,21 @@ We support ALL `volcengine` models, just set `volcengine/<OUR_ENDPOINT_ID>` as a
|
|||
|
||||
```yaml
|
||||
model_list:
|
||||
# Chat model
|
||||
- model_name: volcengine-model
|
||||
litellm_params:
|
||||
model: volcengine/<OUR_ENDPOINT_ID>
|
||||
api_key: os.environ/VOLCENGINE_API_KEY
|
||||
# Embedding model
|
||||
- model_name: volcengine-embedding
|
||||
litellm_params:
|
||||
model: volcengine/doubao-embedding-text-240715
|
||||
api_key: os.environ/VOLCENGINE_API_KEY
|
||||
```
|
||||
|
||||
### Send Request
|
||||
|
||||
#### Chat Completion
|
||||
```shell
|
||||
curl --location 'http://localhost:4000/chat/completions' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
|
|
@ -95,4 +137,15 @@ curl --location 'http://localhost:4000/chat/completions' \
|
|||
}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
#### Embedding
|
||||
```shell
|
||||
curl --location 'http://localhost:4000/embeddings' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"model": "volcengine-embedding",
|
||||
"input": ["hello world", "good morning"]
|
||||
}'
|
||||
```
|
||||
|
|
@ -1215,7 +1215,7 @@ from .llms.jina_ai.embedding.transformation import JinaAIEmbeddingConfig
|
|||
from .llms.xai.chat.transformation import XAIChatConfig
|
||||
from .llms.xai.common_utils import XAIModelInfo
|
||||
from .llms.aiml.chat.transformation import AIMLChatConfig
|
||||
from .llms.volcengine import VolcEngineConfig
|
||||
from .llms.volcengine.chat.transformation import VolcEngineChatConfig as VolcEngineConfig
|
||||
from .llms.codestral.completion.transformation import CodestralTextCompletionConfig
|
||||
from .llms.azure.azure import (
|
||||
AzureOpenAIError,
|
||||
|
|
|
|||
25
litellm/llms/volcengine/__init__.py
Normal file
25
litellm/llms/volcengine/__init__.py
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
"""
|
||||
Volcengine LLM Provider
|
||||
Support for Volcengine (ByteDance) chat and embedding models
|
||||
"""
|
||||
|
||||
from .chat.transformation import VolcEngineChatConfig
|
||||
from .embedding import VolcEngineEmbeddingHandler, VolcEngineEmbeddingConfig
|
||||
from .common_utils import (
|
||||
VolcEngineError,
|
||||
get_volcengine_base_url,
|
||||
get_volcengine_headers,
|
||||
)
|
||||
|
||||
# For backward compatibility, keep the old class name
|
||||
VolcEngineConfig = VolcEngineChatConfig
|
||||
|
||||
__all__ = [
|
||||
"VolcEngineChatConfig",
|
||||
"VolcEngineConfig", # backward compatibility
|
||||
"VolcEngineEmbeddingHandler",
|
||||
"VolcEngineEmbeddingConfig",
|
||||
"VolcEngineError",
|
||||
"get_volcengine_base_url",
|
||||
"get_volcengine_headers",
|
||||
]
|
||||
|
|
@ -3,7 +3,7 @@ from typing import Optional, Union
|
|||
from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
|
||||
|
||||
class VolcEngineConfig(OpenAILikeChatConfig):
|
||||
class VolcEngineChatConfig(OpenAILikeChatConfig):
|
||||
frequency_penalty: Optional[int] = None
|
||||
function_call: Optional[Union[str, dict]] = None
|
||||
functions: Optional[list] = None
|
||||
|
|
@ -82,17 +82,19 @@ class VolcEngineConfig(OpenAILikeChatConfig):
|
|||
|
||||
if "thinking" in optional_params:
|
||||
thinking_value = optional_params.pop("thinking")
|
||||
|
||||
|
||||
# Handle disabled thinking case - don't add to extra_body if disabled
|
||||
if (
|
||||
thinking_value is not None
|
||||
and isinstance(thinking_value, dict)
|
||||
thinking_value is not None
|
||||
and isinstance(thinking_value, dict)
|
||||
and thinking_value.get("type") == "disabled"
|
||||
):
|
||||
# Skip adding thinking parameter when it's disabled
|
||||
pass
|
||||
else:
|
||||
# Add thinking parameter to extra_body for all other cases
|
||||
optional_params.setdefault("extra_body", {})["thinking"] = thinking_value
|
||||
optional_params.setdefault("extra_body", {})[
|
||||
"thinking"
|
||||
] = thinking_value
|
||||
|
||||
return optional_params
|
||||
62
litellm/llms/volcengine/common_utils.py
Normal file
62
litellm/llms/volcengine/common_utils.py
Normal file
|
|
@ -0,0 +1,62 @@
|
|||
"""
|
||||
Common utilities for Volcengine LLM provider
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
|
||||
class VolcEngineError(BaseLLMException):
|
||||
"""
|
||||
Custom exception class for Volcengine provider errors.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, status_code: int, message: str, headers: Optional[httpx.Headers] = None
|
||||
):
|
||||
self.status_code = status_code
|
||||
self.message = message
|
||||
self.headers = headers or httpx.Headers()
|
||||
super().__init__(
|
||||
status_code=status_code, message=message, headers=dict(self.headers)
|
||||
)
|
||||
|
||||
|
||||
def get_volcengine_base_url(api_base: Optional[str] = None) -> str:
|
||||
"""
|
||||
Get the base URL for Volcengine API calls.
|
||||
|
||||
Args:
|
||||
api_base: Optional custom API base URL
|
||||
|
||||
Returns:
|
||||
The base URL to use for API calls
|
||||
"""
|
||||
if api_base:
|
||||
return api_base
|
||||
return "https://ark.cn-beijing.volces.com"
|
||||
|
||||
|
||||
def get_volcengine_headers(api_key: str, extra_headers: Optional[dict] = None) -> dict:
|
||||
"""
|
||||
Get headers for Volcengine API calls.
|
||||
|
||||
Args:
|
||||
api_key: The API key for authentication
|
||||
extra_headers: Optional additional headers
|
||||
|
||||
Returns:
|
||||
Dictionary of headers
|
||||
"""
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
}
|
||||
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
return headers
|
||||
8
litellm/llms/volcengine/embedding/__init__.py
Normal file
8
litellm/llms/volcengine/embedding/__init__.py
Normal file
|
|
@ -0,0 +1,8 @@
|
|||
"""
|
||||
Volcengine Embedding Module
|
||||
"""
|
||||
|
||||
from .handler import VolcEngineEmbeddingHandler
|
||||
from .transformation import VolcEngineEmbeddingConfig
|
||||
|
||||
__all__ = ["VolcEngineEmbeddingHandler", "VolcEngineEmbeddingConfig"]
|
||||
208
litellm/llms/volcengine/embedding/handler.py
Normal file
208
litellm/llms/volcengine/embedding/handler.py
Normal file
|
|
@ -0,0 +1,208 @@
|
|||
"""
|
||||
Volcengine Embedding Handler
|
||||
Handles embedding requests to Volcengine's embedding API
|
||||
"""
|
||||
|
||||
from typing import Dict, List, Optional, Union, Any
|
||||
|
||||
import httpx
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
import litellm
|
||||
|
||||
from .transformation import VolcEngineEmbeddingConfig
|
||||
from ..common_utils import VolcEngineError
|
||||
|
||||
|
||||
class VolcEngineEmbeddingHandler:
|
||||
"""Handler for Volcengine embedding API calls"""
|
||||
|
||||
def __init__(self):
|
||||
self.config = VolcEngineEmbeddingConfig()
|
||||
|
||||
def _convert_to_litellm_response(self, transformed_response: Dict, model: str, input: Union[str, List[str]]) -> EmbeddingResponse:
|
||||
"""Convert transformed response to LiteLLM EmbeddingResponse"""
|
||||
model_response = EmbeddingResponse()
|
||||
model_response.object = transformed_response.get("object", "list")
|
||||
model_response.data = transformed_response.get("data", [])
|
||||
model_response.model = transformed_response.get("model", model)
|
||||
|
||||
# Set usage information
|
||||
usage_data = transformed_response.get("usage", {})
|
||||
if usage_data:
|
||||
model_response.usage = litellm.Usage(
|
||||
prompt_tokens=usage_data.get("prompt_tokens", 0),
|
||||
completion_tokens=0,
|
||||
total_tokens=usage_data.get("total_tokens", usage_data.get("prompt_tokens", 0)),
|
||||
prompt_tokens_details=None,
|
||||
completion_tokens_details=None,
|
||||
)
|
||||
|
||||
return model_response
|
||||
|
||||
def embedding(
|
||||
self,
|
||||
model: str,
|
||||
input: Union[str, List[str]],
|
||||
api_key: str,
|
||||
api_base: Optional[str] = None,
|
||||
encoding_format: Optional[str] = "float",
|
||||
user: Optional[str] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
litellm_logging_obj: Optional[LiteLLMLoggingObj] = None,
|
||||
**kwargs,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
Synchronous embedding call to Volcengine API.
|
||||
|
||||
Args:
|
||||
model: Volcengine model ID (e.g., "doubao-embedding-text-240715")
|
||||
input: Text or list of texts to embed
|
||||
api_key: Volcengine API key
|
||||
api_base: Optional custom API base URL
|
||||
encoding_format: Response format (float, base64, null)
|
||||
user: Optional user identifier
|
||||
timeout: Request timeout
|
||||
extra_headers: Optional additional headers
|
||||
litellm_logging_obj: Optional logging object
|
||||
**kwargs: Additional parameters
|
||||
|
||||
Returns:
|
||||
EmbeddingResponse object
|
||||
"""
|
||||
# Transform request to Volcengine format
|
||||
request_data = self.config.transform_request(
|
||||
model=model,
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
encoding_format=encoding_format,
|
||||
user=user,
|
||||
extra_headers=extra_headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Make HTTP request
|
||||
try:
|
||||
client = HTTPHandler(timeout=timeout)
|
||||
response = client.post(
|
||||
url=request_data["url"],
|
||||
headers=request_data["headers"],
|
||||
json=request_data["data"],
|
||||
)
|
||||
except Exception as e:
|
||||
raise VolcEngineError(
|
||||
status_code=500,
|
||||
message=f"Network error during embedding request: {str(e)}",
|
||||
)
|
||||
|
||||
# Handle HTTP errors
|
||||
if response.status_code != 200:
|
||||
error_message = f"Volcengine embedding request failed with status {response.status_code}"
|
||||
try:
|
||||
error_details = response.json()
|
||||
if "error" in error_details:
|
||||
error_message += f": {error_details['error']}"
|
||||
elif "message" in error_details:
|
||||
error_message += f": {error_details['message']}"
|
||||
except Exception:
|
||||
error_message += f": {response.text}"
|
||||
|
||||
raise VolcEngineError(
|
||||
status_code=response.status_code,
|
||||
message=error_message,
|
||||
headers=response.headers,
|
||||
)
|
||||
|
||||
# Transform response to OpenAI format
|
||||
transformed_response = self.config.transform_response(
|
||||
response=response, model=model, input=input, encoding=encoding_format
|
||||
)
|
||||
|
||||
# Convert to LiteLLM EmbeddingResponse
|
||||
return self._convert_to_litellm_response(transformed_response, model, input)
|
||||
|
||||
async def async_embedding(
|
||||
self,
|
||||
model: str,
|
||||
input: Union[str, List[str]],
|
||||
api_key: str,
|
||||
api_base: Optional[str] = None,
|
||||
encoding_format: Optional[str] = "float",
|
||||
user: Optional[str] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
litellm_logging_obj: Optional[LiteLLMLoggingObj] = None,
|
||||
**kwargs,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
Asynchronous embedding call to Volcengine API.
|
||||
|
||||
Args:
|
||||
model: Volcengine model ID (e.g., "doubao-embedding-text-240715")
|
||||
input: Text or list of texts to embed
|
||||
api_key: Volcengine API key
|
||||
api_base: Optional custom API base URL
|
||||
encoding_format: Response format (float, base64, null)
|
||||
user: Optional user identifier
|
||||
timeout: Request timeout
|
||||
extra_headers: Optional additional headers
|
||||
litellm_logging_obj: Optional logging object
|
||||
**kwargs: Additional parameters
|
||||
|
||||
Returns:
|
||||
EmbeddingResponse object
|
||||
"""
|
||||
# Transform request to Volcengine format
|
||||
request_data = self.config.transform_request(
|
||||
model=model,
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
encoding_format=encoding_format,
|
||||
user=user,
|
||||
extra_headers=extra_headers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# Make async HTTP request
|
||||
try:
|
||||
client = AsyncHTTPHandler(timeout=timeout)
|
||||
response = await client.post(
|
||||
url=request_data["url"],
|
||||
headers=request_data["headers"],
|
||||
json=request_data["data"],
|
||||
)
|
||||
except Exception as e:
|
||||
raise VolcEngineError(
|
||||
status_code=500,
|
||||
message=f"Network error during embedding request: {str(e)}",
|
||||
)
|
||||
|
||||
# Handle HTTP errors
|
||||
if response.status_code != 200:
|
||||
error_message = f"Volcengine embedding request failed with status {response.status_code}"
|
||||
try:
|
||||
error_details = response.json()
|
||||
if "error" in error_details:
|
||||
error_message += f": {error_details['error']}"
|
||||
elif "message" in error_details:
|
||||
error_message += f": {error_details['message']}"
|
||||
except Exception:
|
||||
error_message += f": {response.text}"
|
||||
|
||||
raise VolcEngineError(
|
||||
status_code=response.status_code,
|
||||
message=error_message,
|
||||
headers=response.headers,
|
||||
)
|
||||
|
||||
# Transform response to OpenAI format
|
||||
transformed_response = self.config.transform_response(
|
||||
response=response, model=model, input=input, encoding=encoding_format
|
||||
)
|
||||
|
||||
# Convert to LiteLLM EmbeddingResponse
|
||||
return self._convert_to_litellm_response(transformed_response, model, input)
|
||||
245
litellm/llms/volcengine/embedding/transformation.py
Normal file
245
litellm/llms/volcengine/embedding/transformation.py
Normal file
|
|
@ -0,0 +1,245 @@
|
|||
"""
|
||||
Volcengine Embedding Transformation
|
||||
Transforms OpenAI embedding requests to Volcengine format
|
||||
"""
|
||||
|
||||
from typing import List, Optional, Union, Dict, Any
|
||||
import httpx
|
||||
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from ..common_utils import get_volcengine_base_url, get_volcengine_headers
|
||||
|
||||
|
||||
class VolcEngineEmbeddingConfig(BaseEmbeddingConfig):
|
||||
"""
|
||||
Configuration class for Volcengine embedding models.
|
||||
Reference: https://ark.cn-beijing.volces.com/api/v3/embeddings
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
encoding_format: Optional[str] = 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 super().get_config()
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
"""
|
||||
Get the list of OpenAI parameters supported by Volcengine embedding models.
|
||||
|
||||
Args:
|
||||
model: The model name
|
||||
|
||||
Returns:
|
||||
List of supported parameter names
|
||||
"""
|
||||
return [
|
||||
"encoding_format",
|
||||
"user",
|
||||
"extra_headers",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: Dict[str, Any],
|
||||
optional_params: Dict[str, Any],
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Map OpenAI embedding parameters to Volcengine format.
|
||||
|
||||
Args:
|
||||
non_default_params: Parameters that are not default values
|
||||
optional_params: Optional parameters dict to update
|
||||
model: The model name
|
||||
drop_params: Whether to drop unsupported parameters
|
||||
|
||||
Returns:
|
||||
Updated optional_params dict
|
||||
"""
|
||||
for param, value in non_default_params.items():
|
||||
if param == "encoding_format":
|
||||
# Volcengine supports: float, base64, null
|
||||
if value in ["float", "base64", None]:
|
||||
optional_params["encoding_format"] = value
|
||||
else:
|
||||
if not drop_params:
|
||||
raise ValueError(
|
||||
f"Unsupported encoding_format: {value}. Volcengine supports: float, base64, null"
|
||||
)
|
||||
elif param == "user":
|
||||
# Keep user parameter as-is
|
||||
optional_params["user"] = value
|
||||
elif param in self.get_supported_openai_params(model):
|
||||
optional_params[param] = value
|
||||
elif not drop_params:
|
||||
raise ValueError(f"Unsupported parameter for Volcengine: {param}")
|
||||
|
||||
return optional_params
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
input: Union[str, List[str]],
|
||||
api_key: str,
|
||||
api_base: Optional[str] = None,
|
||||
encoding_format: Optional[str] = "float",
|
||||
user: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
**kwargs,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Transform OpenAI embedding request to Volcengine format.
|
||||
|
||||
Args:
|
||||
model: Model ID (e.g., "doubao-embedding-text-240715")
|
||||
input: Text or list of texts to embed
|
||||
api_key: Volcengine API key
|
||||
api_base: Optional custom API base URL
|
||||
encoding_format: Response format (float, base64, null)
|
||||
user: Optional user identifier
|
||||
extra_headers: Optional additional headers
|
||||
**kwargs: Additional parameters
|
||||
|
||||
Returns:
|
||||
Dict containing url, headers, and data for the request
|
||||
"""
|
||||
# Get base URL
|
||||
base_url = get_volcengine_base_url(api_base)
|
||||
# Avoid duplicate /api/v3 if base_url already contains it
|
||||
if base_url.endswith("/api/v3"):
|
||||
url = f"{base_url}/embeddings"
|
||||
else:
|
||||
url = f"{base_url}/api/v3/embeddings"
|
||||
|
||||
# Get headers
|
||||
headers = get_volcengine_headers(api_key, extra_headers)
|
||||
|
||||
# Prepare request data
|
||||
data = {
|
||||
"model": model,
|
||||
"input": input if isinstance(input, list) else [input],
|
||||
}
|
||||
|
||||
# Add optional parameters
|
||||
if encoding_format is not None:
|
||||
data["encoding_format"] = encoding_format
|
||||
|
||||
return {
|
||||
"url": url,
|
||||
"headers": headers,
|
||||
"data": data,
|
||||
}
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
response: httpx.Response,
|
||||
model: str,
|
||||
input: Union[str, List[str]],
|
||||
encoding: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Transform Volcengine embedding response to OpenAI format.
|
||||
|
||||
Args:
|
||||
response: The HTTP response from Volcengine
|
||||
model: The model used
|
||||
input: The input that was embedded
|
||||
encoding: The encoding format requested
|
||||
|
||||
Returns:
|
||||
OpenAI-compatible embedding response
|
||||
"""
|
||||
try:
|
||||
response_json = response.json()
|
||||
except Exception as e:
|
||||
raise ValueError(f"Failed to parse Volcengine response as JSON: {str(e)}")
|
||||
|
||||
# Volcengine response format matches OpenAI format closely
|
||||
# Just need to ensure all required fields are present
|
||||
transformed_response = {
|
||||
"object": "list",
|
||||
"data": response_json.get("data", []),
|
||||
"model": response_json.get("model", model),
|
||||
"usage": response_json.get("usage", {}),
|
||||
}
|
||||
|
||||
# Add id if present
|
||||
if "id" in response_json:
|
||||
transformed_response["id"] = response_json["id"]
|
||||
|
||||
return transformed_response
|
||||
|
||||
def transform_embedding_request(
|
||||
self,
|
||||
model: str,
|
||||
input: AllEmbeddingInputValues,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""Transform embedding request to Volcengine format"""
|
||||
# Use existing transform_request method
|
||||
return self.transform_request(
|
||||
model=model,
|
||||
input=input,
|
||||
api_key="", # api_key will be in headers
|
||||
**optional_params,
|
||||
)
|
||||
|
||||
def transform_embedding_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: EmbeddingResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str],
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> EmbeddingResponse:
|
||||
"""Transform Volcengine response to EmbeddingResponse"""
|
||||
# Use existing transform_response method
|
||||
transformed_response = self.transform_response(
|
||||
response=raw_response,
|
||||
model=model,
|
||||
input=request_data.get("input", []),
|
||||
)
|
||||
|
||||
# Create EmbeddingResponse from transformed data
|
||||
return EmbeddingResponse(**transformed_response)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""Validate environment and return headers"""
|
||||
# Get Volcengine headers
|
||||
volcengine_headers = get_volcengine_headers(api_key)
|
||||
return {**headers, **volcengine_headers}
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
"""Get error class for Volcengine errors"""
|
||||
from ..common_utils import VolcEngineError
|
||||
return VolcEngineError(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
@ -183,6 +183,7 @@ from .llms.vertex_ai.text_to_speech.text_to_speech_handler import VertexTextToSp
|
|||
from .llms.vertex_ai.vertex_ai_partner_models.main import VertexAIPartnerModels
|
||||
from .llms.vertex_ai.vertex_embeddings.embedding_handler import VertexEmbedding
|
||||
from .llms.vertex_ai.vertex_model_garden.main import VertexAIModelGardenModels
|
||||
from .llms.volcengine.embedding.handler import VolcEngineEmbeddingHandler
|
||||
from .llms.vllm.completion import handler as vllm_handler
|
||||
from .llms.watsonx.chat.handler import WatsonXChatHandler
|
||||
from .llms.watsonx.common_utils import IBMWatsonXMixin
|
||||
|
|
@ -500,7 +501,7 @@ async def acompletion(
|
|||
}
|
||||
if custom_llm_provider is None:
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(
|
||||
model=model, api_base=completion_kwargs.get("base_url", None)
|
||||
model=model, custom_llm_provider=custom_llm_provider, api_base=completion_kwargs.get("base_url", None)
|
||||
)
|
||||
|
||||
fallbacks = fallbacks or litellm.model_fallbacks
|
||||
|
|
@ -3582,7 +3583,7 @@ async def aembedding(*args, **kwargs) -> EmbeddingResponse:
|
|||
model = args[0] if len(args) > 0 else kwargs["model"]
|
||||
### PASS ARGS TO Embedding ###
|
||||
kwargs["aembedding"] = True
|
||||
custom_llm_provider = None
|
||||
custom_llm_provider = kwargs.get("custom_llm_provider", None)
|
||||
try:
|
||||
# Use a partial function to pass your keyword arguments
|
||||
func = partial(embedding, *args, **kwargs)
|
||||
|
|
@ -3592,7 +3593,7 @@ async def aembedding(*args, **kwargs) -> EmbeddingResponse:
|
|||
func_with_context = partial(ctx.run, func)
|
||||
|
||||
_, custom_llm_provider, _, _ = get_llm_provider(
|
||||
model=model, api_base=kwargs.get("api_base", None)
|
||||
model=model, custom_llm_provider=custom_llm_provider, api_base=kwargs.get("api_base", None)
|
||||
)
|
||||
|
||||
# Await normally
|
||||
|
|
@ -4414,6 +4415,46 @@ def embedding( # noqa: PLR0915
|
|||
client=client,
|
||||
aembedding=aembedding,
|
||||
)
|
||||
elif custom_llm_provider == "volcengine":
|
||||
api_key = (
|
||||
api_key
|
||||
or litellm.api_key
|
||||
or get_secret_str("ARK_API_KEY")
|
||||
or get_secret_str("VOLCENGINE_API_KEY")
|
||||
)
|
||||
if api_key is None:
|
||||
raise ValueError(
|
||||
"Missing API key for Volcengine. Set ARK_API_KEY or VOLCENGINE_API_KEY environment variable or pass api_key parameter."
|
||||
)
|
||||
|
||||
handler = VolcEngineEmbeddingHandler()
|
||||
|
||||
if aembedding:
|
||||
response = handler.async_embedding(
|
||||
model=model,
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
encoding_format=optional_params.get("encoding_format", "float"),
|
||||
user=optional_params.get("user"),
|
||||
timeout=timeout,
|
||||
extra_headers=optional_params.get("extra_headers"),
|
||||
litellm_logging_obj=logging,
|
||||
**optional_params,
|
||||
)
|
||||
else:
|
||||
response = handler.embedding(
|
||||
model=model,
|
||||
input=input,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
encoding_format=optional_params.get("encoding_format", "float"),
|
||||
user=optional_params.get("user"),
|
||||
timeout=timeout,
|
||||
extra_headers=optional_params.get("extra_headers"),
|
||||
litellm_logging_obj=logging,
|
||||
**optional_params,
|
||||
)
|
||||
elif custom_llm_provider in litellm._custom_providers:
|
||||
custom_handler: Optional[CustomLLM] = None
|
||||
for item in litellm.custom_provider_map:
|
||||
|
|
|
|||
|
|
@ -5658,6 +5658,11 @@ class Router:
|
|||
)
|
||||
if supported_openai_params is None:
|
||||
supported_openai_params = []
|
||||
|
||||
# Get mode from database model_info if available, otherwise default to "chat"
|
||||
db_model_info = model.get("model_info", {})
|
||||
mode = db_model_info.get("mode", "chat")
|
||||
|
||||
model_info = ModelMapInfo(
|
||||
key=model_group,
|
||||
max_tokens=None,
|
||||
|
|
@ -5666,7 +5671,7 @@ class Router:
|
|||
input_cost_per_token=0,
|
||||
output_cost_per_token=0,
|
||||
litellm_provider=llm_provider,
|
||||
mode="chat",
|
||||
mode=mode,
|
||||
supported_openai_params=supported_openai_params,
|
||||
supports_system_messages=None,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -19746,5 +19746,65 @@
|
|||
"metadata": {
|
||||
"notes": "DALL-E 2 via AI/ML API - Reliable text-to-image generation"
|
||||
}
|
||||
},
|
||||
"doubao-embedding-large": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"output_vector_size": 2048,
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
"litellm_provider": "volcengine",
|
||||
"mode": "embedding",
|
||||
"metadata": {
|
||||
"notes": "Volcengine Doubao embedding model - large version with 2048 dimensions"
|
||||
}
|
||||
},
|
||||
"doubao-embedding-large-text-250515": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"output_vector_size": 2048,
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
"litellm_provider": "volcengine",
|
||||
"mode": "embedding",
|
||||
"metadata": {
|
||||
"notes": "Volcengine Doubao embedding model - text-250515 version with 2048 dimensions"
|
||||
}
|
||||
},
|
||||
"doubao-embedding-large-text-240915": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"output_vector_size": 4096,
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
"litellm_provider": "volcengine",
|
||||
"mode": "embedding",
|
||||
"metadata": {
|
||||
"notes": "Volcengine Doubao embedding model - text-240915 version with 4096 dimensions"
|
||||
}
|
||||
},
|
||||
"doubao-embedding": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"output_vector_size": 2560,
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
"litellm_provider": "volcengine",
|
||||
"mode": "embedding",
|
||||
"metadata": {
|
||||
"notes": "Volcengine Doubao embedding model - standard version with 2560 dimensions"
|
||||
}
|
||||
},
|
||||
"doubao-embedding-text-240715": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"output_vector_size": 2560,
|
||||
"input_cost_per_token": 0.0,
|
||||
"output_cost_per_token": 0.0,
|
||||
"litellm_provider": "volcengine",
|
||||
"mode": "embedding",
|
||||
"metadata": {
|
||||
"notes": "Volcengine Doubao embedding model - text-240715 version with 2560 dimensions"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
262
tests/llm_translation/test_volcengine_embedding.py
Normal file
262
tests/llm_translation/test_volcengine_embedding.py
Normal file
|
|
@ -0,0 +1,262 @@
|
|||
"""
|
||||
Integration tests for Volcengine embedding following LiteLLM testing patterns
|
||||
Based on the BaseLLMEmbeddingTest framework
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
import pytest
|
||||
|
||||
# Add parent directory to path for imports
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
from base_embedding_unit_tests import BaseLLMEmbeddingTest
|
||||
import litellm
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
|
||||
class TestVolcEngineEmbedding(BaseLLMEmbeddingTest):
|
||||
"""Test Volcengine embedding integration following LiteLLM patterns"""
|
||||
|
||||
def get_custom_llm_provider(self) -> litellm.LlmProviders:
|
||||
return litellm.LlmProviders.VOLCENGINE
|
||||
|
||||
def get_base_embedding_call_args(self) -> dict:
|
||||
return {
|
||||
"model": "volcengine/doubao-embedding-text-240715",
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
async def test_basic_embedding(self, sync_mode):
|
||||
"""Test basic embedding functionality with realistic response"""
|
||||
litellm.set_verbose = True
|
||||
embedding_call_args = self.get_base_embedding_call_args()
|
||||
|
||||
# Mock the embedding functions to avoid actual API calls
|
||||
with patch("litellm.embedding") as mock_embedding, patch("litellm.aembedding") as mock_aembedding:
|
||||
# Create realistic Volcengine response
|
||||
mock_response = MagicMock()
|
||||
mock_response.model = "doubao-embedding-text-240715"
|
||||
mock_response.object = "list"
|
||||
mock_response.data = [
|
||||
{
|
||||
"object": "embedding",
|
||||
"embedding": [0.1, 0.2, 0.3] + [0.01 * i for i in range(1021)], # 1024-dim embedding
|
||||
"index": 0
|
||||
},
|
||||
{
|
||||
"object": "embedding",
|
||||
"embedding": [0.4, 0.5, 0.6] + [0.02 * i for i in range(1021)], # 1024-dim embedding
|
||||
"index": 1
|
||||
}
|
||||
]
|
||||
mock_response.usage.prompt_tokens = 2
|
||||
mock_response.usage.total_tokens = 2
|
||||
|
||||
mock_embedding.return_value = mock_response
|
||||
mock_aembedding.return_value = mock_response
|
||||
|
||||
# Test sync mode
|
||||
if sync_mode is True:
|
||||
response = litellm.embedding(
|
||||
**embedding_call_args,
|
||||
input=["hello", "world"],
|
||||
)
|
||||
|
||||
# Verify response structure matches Volcengine format
|
||||
assert response.model == "doubao-embedding-text-240715"
|
||||
assert response.object == "list"
|
||||
assert len(response.data) == 2
|
||||
assert len(response.data[0]["embedding"]) == 1024
|
||||
assert response.usage.total_tokens > 0
|
||||
|
||||
# Test async mode
|
||||
else:
|
||||
response = await litellm.aembedding(
|
||||
**embedding_call_args,
|
||||
input=["hello", "world"],
|
||||
)
|
||||
|
||||
# Verify response structure
|
||||
assert response.model == "doubao-embedding-text-240715"
|
||||
assert response.object == "list"
|
||||
assert len(response.data) == 2
|
||||
assert len(response.data[0]["embedding"]) == 1024
|
||||
assert response.usage.total_tokens > 0
|
||||
|
||||
|
||||
def test_volcengine_embedding_with_encoding_formats():
|
||||
"""Test Volcengine embedding with different encoding formats"""
|
||||
|
||||
test_cases = [
|
||||
{"encoding_format": "float"},
|
||||
{"encoding_format": "base64"},
|
||||
{"encoding_format": None}, # Default
|
||||
]
|
||||
|
||||
for params in test_cases:
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
# Create mock response based on encoding format
|
||||
mock_response = MagicMock()
|
||||
mock_response.model = "doubao-embedding-text-240715"
|
||||
mock_response.object = "list"
|
||||
|
||||
if params["encoding_format"] == "base64":
|
||||
# Simulate base64 encoded embeddings
|
||||
mock_response.data = [
|
||||
{
|
||||
"object": "embedding",
|
||||
"embedding": "c29tZS1iYXNlNjQtZW5jb2RlZC1lbWJlZGRpbmc=", # base64 encoded
|
||||
"index": 0
|
||||
}
|
||||
]
|
||||
else:
|
||||
# Float embeddings (default)
|
||||
mock_response.data = [
|
||||
{
|
||||
"object": "embedding",
|
||||
"embedding": [0.1, 0.2, 0.3, -0.1] * 256, # 1024 dimensions
|
||||
"index": 0
|
||||
}
|
||||
]
|
||||
|
||||
mock_response.usage.prompt_tokens = 3
|
||||
mock_response.usage.total_tokens = 3
|
||||
mock_embedding.return_value = mock_response
|
||||
|
||||
# Test the call
|
||||
litellm.embedding(
|
||||
model="volcengine/doubao-embedding-text-240715",
|
||||
input=["test text"],
|
||||
**params
|
||||
)
|
||||
|
||||
# Verify the call was made with correct parameters
|
||||
mock_embedding.assert_called_once()
|
||||
call_args = mock_embedding.call_args
|
||||
assert call_args[1]["model"] == "volcengine/doubao-embedding-text-240715"
|
||||
assert call_args[1]["input"] == ["test text"]
|
||||
|
||||
if params["encoding_format"] is not None:
|
||||
assert call_args[1]["encoding_format"] == params["encoding_format"]
|
||||
|
||||
|
||||
def test_volcengine_embedding_with_user_parameter():
|
||||
"""Test Volcengine embedding with user parameter for tracking"""
|
||||
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_response = MagicMock()
|
||||
mock_response.model = "doubao-embedding-text-240715"
|
||||
mock_response.object = "list"
|
||||
mock_response.data = [
|
||||
{
|
||||
"object": "embedding",
|
||||
"embedding": [0.1] * 1024,
|
||||
"index": 0
|
||||
}
|
||||
]
|
||||
mock_response.usage.prompt_tokens = 5
|
||||
mock_response.usage.total_tokens = 5
|
||||
mock_embedding.return_value = mock_response
|
||||
|
||||
# Test with user parameter
|
||||
litellm.embedding(
|
||||
model="volcengine/doubao-embedding-text-240715",
|
||||
input=["user tracking test"],
|
||||
user="test-user-12345"
|
||||
)
|
||||
|
||||
# Verify user parameter was passed
|
||||
mock_embedding.assert_called_once()
|
||||
call_args = mock_embedding.call_args
|
||||
assert call_args[1]["user"] == "test-user-12345"
|
||||
|
||||
|
||||
def test_volcengine_embedding_error_scenarios():
|
||||
"""Test Volcengine embedding error handling in integration context"""
|
||||
|
||||
error_scenarios = [
|
||||
# Invalid model name
|
||||
{
|
||||
"model": "volcengine/invalid-model-name",
|
||||
"expected_error_pattern": "model"
|
||||
},
|
||||
# Invalid encoding format
|
||||
{
|
||||
"model": "volcengine/doubao-embedding-text-240715",
|
||||
"encoding_format": "invalid_format",
|
||||
"expected_error_pattern": "encoding_format"
|
||||
}
|
||||
]
|
||||
|
||||
for scenario in error_scenarios:
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
# Configure mock to raise appropriate errors
|
||||
if "invalid-model" in scenario.get("model", ""):
|
||||
mock_embedding.side_effect = Exception("Model not found")
|
||||
elif scenario.get("encoding_format") == "invalid_format":
|
||||
mock_embedding.side_effect = ValueError("Unsupported encoding_format")
|
||||
|
||||
# Test that errors are properly raised
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
test_params = {k: v for k, v in scenario.items() if k != "expected_error_pattern"}
|
||||
litellm.embedding(
|
||||
input=["test"],
|
||||
**test_params
|
||||
)
|
||||
|
||||
# Verify error message contains expected pattern
|
||||
assert scenario["expected_error_pattern"].lower() in str(exc_info.value).lower()
|
||||
|
||||
|
||||
def test_volcengine_embedding_with_multiple_inputs():
|
||||
"""Test Volcengine embedding with various input lengths and types"""
|
||||
|
||||
test_inputs = [
|
||||
# Single short text
|
||||
["hello"],
|
||||
# Multiple short texts
|
||||
["hello", "world", "test"],
|
||||
# Mixed length texts
|
||||
["short", "This is a much longer text that should be handled properly by the embedding service"],
|
||||
# Unicode content
|
||||
["测试中文文本", "Test English text", "混合语言 mixed language"],
|
||||
# Many inputs (batch processing)
|
||||
[f"Test sentence number {i}" for i in range(10)]
|
||||
]
|
||||
|
||||
for test_input in test_inputs:
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
# Create proportional mock response
|
||||
mock_response = MagicMock()
|
||||
mock_response.model = "doubao-embedding-text-240715"
|
||||
mock_response.object = "list"
|
||||
mock_response.data = [
|
||||
{
|
||||
"object": "embedding",
|
||||
"embedding": [0.1 * (i + 1)] * 1024, # Unique embedding per input
|
||||
"index": i
|
||||
}
|
||||
for i in range(len(test_input))
|
||||
]
|
||||
mock_response.usage.prompt_tokens = len(test_input) * 5 # Realistic token estimate
|
||||
mock_response.usage.total_tokens = len(test_input) * 5
|
||||
mock_embedding.return_value = mock_response
|
||||
|
||||
# Test the call
|
||||
response = litellm.embedding(
|
||||
model="volcengine/doubao-embedding-text-240715",
|
||||
input=test_input
|
||||
)
|
||||
|
||||
# Verify response matches input count
|
||||
assert len(response.data) == len(test_input)
|
||||
for i, embedding_data in enumerate(response.data):
|
||||
assert embedding_data["index"] == i
|
||||
assert len(embedding_data["embedding"]) == 1024
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
1
tests/test_litellm/llms/volcengine/__init__.py
Normal file
1
tests/test_litellm/llms/volcengine/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
# Volcengine tests
|
||||
1
tests/test_litellm/llms/volcengine/embedding/__init__.py
Normal file
1
tests/test_litellm/llms/volcengine/embedding/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
# Volcengine embedding tests
|
||||
|
|
@ -0,0 +1,450 @@
|
|||
"""
|
||||
Improved tests for Volcengine Embedding functionality
|
||||
Tests real business logic without excessive mocking
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import json
|
||||
import httpx
|
||||
from unittest.mock import Mock, patch, MagicMock
|
||||
from typing import List, Dict, Any
|
||||
|
||||
from litellm.llms.volcengine.embedding import VolcEngineEmbeddingHandler, VolcEngineEmbeddingConfig
|
||||
from litellm.llms.volcengine.common_utils import VolcEngineError
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.types.llms.openai import AllEmbeddingInputValues
|
||||
|
||||
|
||||
class TestVolcEngineEmbeddingConfigBusinessLogic:
|
||||
"""Test real business logic of VolcEngineEmbeddingConfig without excessive mocking"""
|
||||
|
||||
def setup_method(self):
|
||||
"""Setup test fixtures"""
|
||||
self.config = VolcEngineEmbeddingConfig()
|
||||
self.model = "doubao-embedding-text-240715"
|
||||
self.api_key = "test-api-key-12345"
|
||||
|
||||
def test_supported_params_completeness(self):
|
||||
"""Test that all required parameters are supported"""
|
||||
params = self.config.get_supported_openai_params(self.model)
|
||||
|
||||
# Verify essential parameters are supported
|
||||
required_params = ["encoding_format", "user", "extra_headers"]
|
||||
for param in required_params:
|
||||
assert param in params, f"Required parameter '{param}' not supported"
|
||||
|
||||
def test_parameter_mapping_with_valid_values(self):
|
||||
"""Test parameter mapping with various valid values"""
|
||||
test_cases = [
|
||||
# Standard float encoding
|
||||
{"encoding_format": "float", "user": "test-user"},
|
||||
# Base64 encoding
|
||||
{"encoding_format": "base64", "user": "batch-user"},
|
||||
# None encoding (default)
|
||||
{"encoding_format": None, "user": "api-user"},
|
||||
# Only user parameter
|
||||
{"user": "minimal-user"},
|
||||
]
|
||||
|
||||
for test_params in test_cases:
|
||||
result = self.config.map_openai_params(
|
||||
non_default_params=test_params,
|
||||
optional_params={},
|
||||
model=self.model,
|
||||
drop_params=False
|
||||
)
|
||||
|
||||
# Verify all valid parameters are preserved
|
||||
for key, value in test_params.items():
|
||||
if value is not None:
|
||||
assert result[key] == value, f"Parameter {key} not mapped correctly"
|
||||
|
||||
def test_parameter_mapping_with_invalid_encoding(self):
|
||||
"""Test proper error handling for invalid encoding formats"""
|
||||
invalid_encodings = ["int32", "binary", "invalid_format", 123, []]
|
||||
|
||||
for invalid_encoding in invalid_encodings:
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
self.config.map_openai_params(
|
||||
non_default_params={"encoding_format": invalid_encoding},
|
||||
optional_params={},
|
||||
model=self.model,
|
||||
drop_params=False
|
||||
)
|
||||
|
||||
assert "Unsupported encoding_format" in str(exc_info.value)
|
||||
assert str(invalid_encoding) in str(exc_info.value)
|
||||
|
||||
def test_parameter_dropping_behavior(self):
|
||||
"""Test parameter dropping when drop_params=True"""
|
||||
invalid_params = {
|
||||
"encoding_format": "invalid_format",
|
||||
"unsupported_param": "value",
|
||||
"another_invalid": 123
|
||||
}
|
||||
|
||||
result = self.config.map_openai_params(
|
||||
non_default_params=invalid_params,
|
||||
optional_params={},
|
||||
model=self.model,
|
||||
drop_params=True
|
||||
)
|
||||
|
||||
# Should drop all invalid parameters
|
||||
for param in invalid_params.keys():
|
||||
assert param not in result, f"Invalid parameter {param} was not dropped"
|
||||
|
||||
def test_request_transformation_structure(self):
|
||||
"""Test request transformation produces correct structure"""
|
||||
test_inputs = [
|
||||
# Single string input
|
||||
"Hello world",
|
||||
# Multiple strings
|
||||
["Hello", "World", "Test"],
|
||||
# Mixed content
|
||||
["Short", "This is a longer text for testing purposes"],
|
||||
]
|
||||
|
||||
for input_data in test_inputs:
|
||||
result = self.config.transform_request(
|
||||
model=self.model,
|
||||
input=input_data,
|
||||
api_key=self.api_key,
|
||||
encoding_format="float"
|
||||
)
|
||||
|
||||
# Verify structure
|
||||
assert "url" in result
|
||||
assert "headers" in result
|
||||
assert "data" in result
|
||||
|
||||
# Verify URL
|
||||
assert result["url"] == "https://ark.cn-beijing.volces.com/api/v3/embeddings"
|
||||
|
||||
# Verify headers
|
||||
headers = result["headers"]
|
||||
assert headers["Authorization"] == f"Bearer {self.api_key}"
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
|
||||
# Verify data
|
||||
data = result["data"]
|
||||
assert data["model"] == self.model
|
||||
assert data["encoding_format"] == "float"
|
||||
|
||||
# Input should always be a list
|
||||
if isinstance(input_data, str):
|
||||
assert data["input"] == [input_data]
|
||||
else:
|
||||
assert data["input"] == input_data
|
||||
|
||||
def test_response_transformation_with_real_data(self):
|
||||
"""Test response transformation with realistic Volcengine response data"""
|
||||
# Simulate real Volcengine API response
|
||||
volcengine_responses = [
|
||||
# Single embedding response
|
||||
{
|
||||
"id": "cmpl-123456789",
|
||||
"object": "list",
|
||||
"model": "doubao-embedding-text-240715",
|
||||
"data": [
|
||||
{
|
||||
"object": "embedding",
|
||||
"index": 0,
|
||||
"embedding": [0.1, -0.2, 0.3, 0.4, -0.5] * 100 # Realistic embedding size
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 5,
|
||||
"total_tokens": 5
|
||||
}
|
||||
},
|
||||
# Multiple embeddings response
|
||||
{
|
||||
"id": "cmpl-987654321",
|
||||
"object": "list",
|
||||
"model": "doubao-embedding-text-240715",
|
||||
"data": [
|
||||
{
|
||||
"object": "embedding",
|
||||
"index": 0,
|
||||
"embedding": [0.1, 0.2, 0.3] * 256
|
||||
},
|
||||
{
|
||||
"object": "embedding",
|
||||
"index": 1,
|
||||
"embedding": [0.4, 0.5, 0.6] * 256
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 12,
|
||||
"total_tokens": 12
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
for response_data in volcengine_responses:
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.json.return_value = response_data
|
||||
|
||||
result = self.config.transform_response(
|
||||
response=mock_response,
|
||||
model=self.model,
|
||||
input=["test input"],
|
||||
)
|
||||
|
||||
# Verify transformation preserves important data
|
||||
assert result["object"] == "list"
|
||||
assert result["model"] == response_data["model"]
|
||||
assert len(result["data"]) == len(response_data["data"])
|
||||
assert result["usage"] == response_data["usage"]
|
||||
|
||||
# Verify embedding data integrity
|
||||
for i, embedding_item in enumerate(result["data"]):
|
||||
original_item = response_data["data"][i]
|
||||
assert embedding_item["object"] == "embedding"
|
||||
assert embedding_item["index"] == original_item["index"]
|
||||
assert len(embedding_item["embedding"]) == len(original_item["embedding"])
|
||||
|
||||
def test_response_transformation_with_error_data(self):
|
||||
"""Test response transformation handles error response formats correctly"""
|
||||
# Test that transform_response can handle both success and error response structures
|
||||
|
||||
# Success response (should work)
|
||||
success_response = {
|
||||
"id": "cmpl-123",
|
||||
"object": "list",
|
||||
"model": "doubao-embedding-text-240715",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}],
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2}
|
||||
}
|
||||
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.json.return_value = success_response
|
||||
|
||||
result = self.config.transform_response(
|
||||
response=mock_response,
|
||||
model=self.model,
|
||||
input=["test"],
|
||||
)
|
||||
|
||||
# Should successfully transform
|
||||
assert result["object"] == "list"
|
||||
assert result["model"] == "doubao-embedding-text-240715"
|
||||
|
||||
# Error response (should still transform but with empty/missing data)
|
||||
error_response = {
|
||||
"error": {
|
||||
"message": "Rate limit exceeded",
|
||||
"type": "rate_limit_error"
|
||||
}
|
||||
}
|
||||
|
||||
mock_response.json.return_value = error_response
|
||||
|
||||
result = self.config.transform_response(
|
||||
response=mock_response,
|
||||
model=self.model,
|
||||
input=["test"],
|
||||
)
|
||||
|
||||
# Should handle missing fields gracefully
|
||||
assert result["object"] == "list" # default value
|
||||
assert result["data"] == [] # default empty data
|
||||
assert result["usage"] == {} # default empty usage
|
||||
|
||||
|
||||
class TestVolcEngineEmbeddingHandlerBusinessLogic:
|
||||
"""Test VolcEngineEmbeddingHandler with focus on business logic"""
|
||||
|
||||
def setup_method(self):
|
||||
self.handler = VolcEngineEmbeddingHandler()
|
||||
self.model = "doubao-embedding-text-240715"
|
||||
self.api_key = "test-api-key-12345"
|
||||
|
||||
def test_response_conversion_to_litellm_format(self):
|
||||
"""Test conversion of Volcengine response to LiteLLM EmbeddingResponse"""
|
||||
volcengine_response = {
|
||||
"id": "emb-123",
|
||||
"object": "list",
|
||||
"model": self.model,
|
||||
"data": [
|
||||
{
|
||||
"object": "embedding",
|
||||
"index": 0,
|
||||
"embedding": [0.1, 0.2, 0.3, -0.1, -0.2] * 200 # 1000-dimensional embedding
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 8,
|
||||
"total_tokens": 8
|
||||
}
|
||||
}
|
||||
|
||||
result = self.handler._convert_to_litellm_response(
|
||||
volcengine_response,
|
||||
self.model,
|
||||
["test input"]
|
||||
)
|
||||
|
||||
# Verify result is proper EmbeddingResponse
|
||||
assert isinstance(result, EmbeddingResponse)
|
||||
assert result.object == "list"
|
||||
assert result.model == self.model
|
||||
assert len(result.data) == 1
|
||||
assert len(result.data[0]["embedding"]) == 1000
|
||||
|
||||
# Verify usage information
|
||||
assert result.usage.prompt_tokens == 8
|
||||
assert result.usage.total_tokens == 8
|
||||
assert result.usage.completion_tokens == 0
|
||||
|
||||
def test_network_error_handling_without_mocking_business_logic(self):
|
||||
"""Test network error handling preserves business logic"""
|
||||
|
||||
# Test with actual VolcEngineError class
|
||||
with pytest.raises(VolcEngineError) as exc_info:
|
||||
# This would raise a network error in real scenario
|
||||
error = VolcEngineError(
|
||||
status_code=500,
|
||||
message="Network error during embedding request: Connection timeout"
|
||||
)
|
||||
raise error
|
||||
|
||||
# Verify error contains meaningful information
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "Network error during embedding request" in str(exc_info.value.message)
|
||||
assert "Connection timeout" in str(exc_info.value.message)
|
||||
|
||||
def test_input_validation_and_preprocessing(self):
|
||||
"""Test input validation and preprocessing logic"""
|
||||
test_cases = [
|
||||
# String input should be converted to list
|
||||
("single string", ["single string"]),
|
||||
# List input should remain list
|
||||
(["multiple", "strings"], ["multiple", "strings"]),
|
||||
# Empty string handling
|
||||
("", [""]),
|
||||
# Unicode handling
|
||||
("测试中文", ["测试中文"]),
|
||||
# Special characters
|
||||
("Special chars: @#$%^&*()", ["Special chars: @#$%^&*()"]),
|
||||
]
|
||||
|
||||
for input_data, expected_output in test_cases:
|
||||
# Test the actual transformation logic
|
||||
config = VolcEngineEmbeddingConfig()
|
||||
result = config.transform_request(
|
||||
model=self.model,
|
||||
input=input_data,
|
||||
api_key=self.api_key,
|
||||
)
|
||||
|
||||
assert result["data"]["input"] == expected_output
|
||||
|
||||
|
||||
class TestVolcEngineEmbeddingIntegration:
|
||||
"""Integration tests that test the full pipeline with minimal mocking"""
|
||||
|
||||
def setup_method(self):
|
||||
self.handler = VolcEngineEmbeddingHandler()
|
||||
self.model = "doubao-embedding-text-240715"
|
||||
self.api_key = "test-api-key-12345"
|
||||
|
||||
def test_full_request_response_cycle(self):
|
||||
"""Test the complete request-response cycle with realistic data"""
|
||||
|
||||
# Create a realistic Volcengine response
|
||||
realistic_response_data = {
|
||||
"id": "cmpl-uqkvlQyYK7bGYrRHQ0eXlWi6",
|
||||
"object": "list",
|
||||
"model": "doubao-embedding-text-240715",
|
||||
"data": [
|
||||
{
|
||||
"object": "embedding",
|
||||
"index": 0,
|
||||
"embedding": [0.0023064255] + [0.1 * (i % 10 - 5) for i in range(1023)] # Realistic 1024-dim embedding
|
||||
},
|
||||
{
|
||||
"object": "embedding",
|
||||
"index": 1,
|
||||
"embedding": [-0.0038562391] + [0.05 * (i % 20 - 10) for i in range(1023)]
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 6,
|
||||
"total_tokens": 6
|
||||
}
|
||||
}
|
||||
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = realistic_response_data
|
||||
|
||||
# Only mock the HTTP call, not the business logic
|
||||
with patch('litellm.llms.volcengine.embedding.handler.HTTPHandler') as mock_handler:
|
||||
mock_client = Mock()
|
||||
mock_client.post.return_value = mock_response
|
||||
mock_handler.return_value = mock_client
|
||||
|
||||
# Test the actual embedding call
|
||||
result = self.handler.embedding(
|
||||
model=self.model,
|
||||
input=["Hello world", "Test embedding"],
|
||||
api_key=self.api_key,
|
||||
encoding_format="float"
|
||||
)
|
||||
|
||||
# Verify the HTTP request was made correctly (this tests integration)
|
||||
mock_client.post.assert_called_once()
|
||||
call_args = mock_client.post.call_args
|
||||
|
||||
# Verify request structure
|
||||
assert call_args.kwargs["url"] == "https://ark.cn-beijing.volces.com/api/v3/embeddings"
|
||||
assert call_args.kwargs["headers"]["Authorization"] == f"Bearer {self.api_key}"
|
||||
|
||||
request_data = call_args.kwargs["json"]
|
||||
assert request_data["model"] == self.model
|
||||
assert request_data["input"] == ["Hello world", "Test embedding"]
|
||||
assert request_data["encoding_format"] == "float"
|
||||
|
||||
# Verify the response processing (real business logic)
|
||||
assert isinstance(result, EmbeddingResponse)
|
||||
assert result.model == self.model
|
||||
assert len(result.data) == 2
|
||||
assert len(result.data[0]["embedding"]) == 1024
|
||||
assert len(result.data[1]["embedding"]) == 1024
|
||||
assert result.usage.prompt_tokens == 6
|
||||
|
||||
def test_parameter_validation_integration(self):
|
||||
"""Test parameter validation in the full integration context"""
|
||||
|
||||
# Test with various parameter combinations that should work
|
||||
valid_param_sets = [
|
||||
{"encoding_format": "float"},
|
||||
{"encoding_format": "base64"},
|
||||
{"user": "test-user-123"},
|
||||
{"encoding_format": "float", "user": "test-user"},
|
||||
{"extra_headers": {"Custom-Header": "value"}},
|
||||
]
|
||||
|
||||
for params in valid_param_sets:
|
||||
# Only create the request, don't execute (avoids HTTP call)
|
||||
config = VolcEngineEmbeddingConfig()
|
||||
try:
|
||||
result = config.transform_request(
|
||||
model=self.model,
|
||||
input=["test"],
|
||||
api_key=self.api_key,
|
||||
**params
|
||||
)
|
||||
# Verify structure is correct
|
||||
assert "url" in result
|
||||
assert "headers" in result
|
||||
assert "data" in result
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(f"Valid parameters {params} caused error: {e}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
|
@ -4,7 +4,7 @@ from unittest.mock import MagicMock, patch
|
|||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.llms.volcengine import VolcEngineConfig
|
||||
from litellm.llms.volcengine.chat.transformation import VolcEngineChatConfig as VolcEngineConfig
|
||||
from litellm.utils import get_optional_params
|
||||
|
||||
|
||||
|
|
@ -170,7 +170,7 @@ def test_all_model_configs():
|
|||
drop_params=False,
|
||||
) == {"max_tokens": 10}
|
||||
|
||||
from litellm.llms.volcengine import VolcEngineConfig
|
||||
from litellm.llms.volcengine.chat.transformation import VolcEngineChatConfig as VolcEngineConfig
|
||||
|
||||
assert "max_completion_tokens" in VolcEngineConfig().get_supported_openai_params(
|
||||
model="llama3"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue