mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
[Feat] Allow adding OpenAI compatible chat providers using .json + add public ai provider (#17448)
* feat: Add JSON config for OpenAI-compatible providers Co-authored-by: ishaan <ishaan@berri.ai> * feat: Add simple JSON config for OpenAI-compatible providers Co-authored-by: ishaan <ishaan@berri.ai> * feat: Implement JSON-based provider config and migrate PublicAI Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * Checkpoint before follow-up message Co-authored-by: ishaan <ishaan@berri.ai> * docs fix * undo change --------- Co-authored-by: Cursor Agent <cursoragent@cursor.com> Co-authored-by: ishaan <ishaan@berri.ai>
This commit is contained in:
parent
fadfbb13d3
commit
b2e8d3fd42
13 changed files with 827 additions and 125 deletions
|
|
@ -0,0 +1,130 @@
|
|||
# Adding OpenAI-Compatible Providers
|
||||
|
||||
For simple OpenAI-compatible providers (like Hyperbolic, Nscale, etc.), you can add support by editing a single JSON file.
|
||||
|
||||
## Quick Start
|
||||
|
||||
1. Edit `litellm/llms/openai_like/providers.json`
|
||||
2. Add your provider configuration
|
||||
3. Test with: `litellm.completion(model="your_provider/model-name", ...)`
|
||||
|
||||
## Basic Configuration
|
||||
|
||||
For a fully OpenAI-compatible provider:
|
||||
|
||||
```json
|
||||
{
|
||||
"your_provider": {
|
||||
"base_url": "https://api.yourprovider.com/v1",
|
||||
"api_key_env": "YOUR_PROVIDER_API_KEY"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
That's it! The provider is now available.
|
||||
|
||||
## Configuration Options
|
||||
|
||||
### Required Fields
|
||||
|
||||
- `base_url` - API endpoint (e.g., `https://api.provider.com/v1`)
|
||||
- `api_key_env` - Environment variable name for API key (e.g., `PROVIDER_API_KEY`)
|
||||
|
||||
### Optional Fields
|
||||
|
||||
- `api_base_env` - Environment variable to override `base_url`
|
||||
- `base_class` - Use `"openai_gpt"` (default) or `"openai_like"`
|
||||
- `param_mappings` - Map OpenAI parameter names to provider-specific names
|
||||
- `constraints` - Parameter value constraints (min/max)
|
||||
- `special_handling` - Special behaviors like content format conversion
|
||||
|
||||
## Examples
|
||||
|
||||
### Simple Provider (Fully Compatible)
|
||||
|
||||
```json
|
||||
{
|
||||
"hyperbolic": {
|
||||
"base_url": "https://api.hyperbolic.xyz/v1",
|
||||
"api_key_env": "HYPERBOLIC_API_KEY"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Provider with Parameter Mapping
|
||||
|
||||
```json
|
||||
{
|
||||
"publicai": {
|
||||
"base_url": "https://api.publicai.co/v1",
|
||||
"api_key_env": "PUBLICAI_API_KEY",
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Provider with Constraints
|
||||
|
||||
```json
|
||||
{
|
||||
"custom_provider": {
|
||||
"base_url": "https://api.custom.com/v1",
|
||||
"api_key_env": "CUSTOM_API_KEY",
|
||||
"constraints": {
|
||||
"temperature_max": 1.0,
|
||||
"temperature_min": 0.0
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
import litellm
|
||||
import os
|
||||
|
||||
# Set your API key
|
||||
os.environ["YOUR_PROVIDER_API_KEY"] = "your-key-here"
|
||||
|
||||
# Use the provider
|
||||
response = litellm.completion(
|
||||
model="your_provider/model-name",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
)
|
||||
```
|
||||
|
||||
## When to Use Python Instead
|
||||
|
||||
Use a Python config class if you need:
|
||||
|
||||
- Custom authentication flows (OAuth, JWT, etc.)
|
||||
- Complex request/response transformations
|
||||
- Provider-specific streaming logic
|
||||
- Advanced tool calling modifications
|
||||
|
||||
For these cases, create a config class in `litellm/llms/your_provider/chat/transformation.py` that inherits from `OpenAIGPTConfig` or `OpenAILikeChatConfig`.
|
||||
|
||||
## Testing
|
||||
|
||||
Test your provider:
|
||||
|
||||
```bash
|
||||
# Quick test
|
||||
python -c "
|
||||
import litellm
|
||||
import os
|
||||
os.environ['PROVIDER_API_KEY'] = 'your-key'
|
||||
response = litellm.completion(
|
||||
model='provider/model-name',
|
||||
messages=[{'role': 'user', 'content': 'test'}]
|
||||
)
|
||||
print(response.choices[0].message.content)
|
||||
"
|
||||
```
|
||||
|
||||
## Reference
|
||||
|
||||
See existing providers in `litellm/llms/openai_like/providers.json` for examples.
|
||||
|
|
@ -11,7 +11,7 @@ Selecting `openai` as the provider routes your request to an OpenAI-compatible e
|
|||
This library **requires** an API key for all requests, either through the `api_key` parameter
|
||||
or the `OPENAI_API_KEY` environment variable.
|
||||
|
||||
If you don’t want to provide a fake API key in each request, consider using a provider that directly matches your
|
||||
If you don't want to provide a fake API key in each request, consider using a provider that directly matches your
|
||||
OpenAI-compatible endpoint, such as [`hosted_vllm`](/docs/providers/vllm) or [`llamafile`](/docs/providers/llamafile).
|
||||
|
||||
:::
|
||||
|
|
@ -150,4 +150,4 @@ model_list:
|
|||
api_base: http://my-custom-base
|
||||
api_key: ""
|
||||
supports_system_message: False # 👈 KEY CHANGE
|
||||
```
|
||||
```
|
||||
|
|
|
|||
|
|
@ -1339,7 +1339,7 @@ from .llms.nebius.chat.transformation import NebiusConfig
|
|||
from .llms.wandb.chat.transformation import WandbConfig
|
||||
from .llms.dashscope.chat.transformation import DashScopeChatConfig
|
||||
from .llms.moonshot.chat.transformation import MoonshotChatConfig
|
||||
from .llms.publicai.chat.transformation import PublicAIChatConfig
|
||||
# PublicAI now uses JSON-based configuration (see litellm/llms/openai_like/providers.json)
|
||||
from .llms.docker_model_runner.chat.transformation import DockerModelRunnerChatConfig
|
||||
from .llms.v0.chat.transformation import V0ChatConfig
|
||||
from .llms.oci.chat.transformation import OCIChatConfig
|
||||
|
|
|
|||
|
|
@ -545,7 +545,7 @@ openai_compatible_endpoints: List = [
|
|||
"api.studio.nebius.ai/v1",
|
||||
"https://dashscope-intl.aliyuncs.com/compatible-mode/v1",
|
||||
"https://api.moonshot.ai/v1",
|
||||
"https://platform.publicai.co/v1",
|
||||
"https://api.publicai.co/v1",
|
||||
"https://api.v0.dev/v1",
|
||||
"https://api.morphllm.com/v1",
|
||||
"https://api.lambda.ai/v1",
|
||||
|
|
@ -587,6 +587,7 @@ openai_compatible_providers: List = [
|
|||
"github_copilot", # GitHub Copilot Chat API
|
||||
"novita",
|
||||
"meta_llama",
|
||||
"publicai", # PublicAI - JSON-configured provider
|
||||
"featherless_ai",
|
||||
"nscale",
|
||||
"nebius",
|
||||
|
|
|
|||
|
|
@ -468,6 +468,18 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
|
|||
custom_llm_provider = model.split("/", 1)[0]
|
||||
model = model.split("/", 1)[1]
|
||||
|
||||
# Check JSON providers FIRST (before hardcoded ones)
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
|
||||
if JSONProviderRegistry.exists(custom_llm_provider):
|
||||
provider_config = JSONProviderRegistry.get(custom_llm_provider)
|
||||
config_class = create_config_class(provider_config)
|
||||
api_base, dynamic_api_key = config_class()._get_openai_compatible_provider_info(
|
||||
api_base, api_key
|
||||
)
|
||||
return model, custom_llm_provider, dynamic_api_key, api_base
|
||||
|
||||
if custom_llm_provider == "perplexity":
|
||||
# perplexity is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.perplexity.ai
|
||||
(
|
||||
|
|
@ -763,13 +775,7 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
|
|||
) = litellm.MoonshotChatConfig()._get_openai_compatible_provider_info(
|
||||
api_base, api_key
|
||||
)
|
||||
elif custom_llm_provider == "publicai":
|
||||
(
|
||||
api_base,
|
||||
dynamic_api_key,
|
||||
) = litellm.PublicAIChatConfig()._get_openai_compatible_provider_info(
|
||||
api_base, api_key
|
||||
)
|
||||
# publicai is now handled by JSON config (see litellm/llms/openai_like/providers.json)
|
||||
elif custom_llm_provider == "docker_model_runner":
|
||||
(
|
||||
api_base,
|
||||
|
|
|
|||
129
litellm/llms/openai_like/README.md
Normal file
129
litellm/llms/openai_like/README.md
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
# JSON-Based OpenAI-Compatible Provider Configuration
|
||||
|
||||
This directory contains the new JSON-based configuration system for OpenAI-compatible providers.
|
||||
|
||||
## Overview
|
||||
|
||||
Instead of creating a full Python module for simple OpenAI-compatible providers, you can now define them in a single JSON file.
|
||||
|
||||
## Files
|
||||
|
||||
- `providers.json` - Configuration file for all JSON-based providers
|
||||
- `json_loader.py` - Loads and parses the JSON configuration
|
||||
- `dynamic_config.py` - Generates Python config classes from JSON
|
||||
- `chat/` - Existing OpenAI-like chat completion handlers
|
||||
|
||||
## Adding a New Provider
|
||||
|
||||
### For Simple OpenAI-Compatible Providers
|
||||
|
||||
Edit `providers.json` and add your provider:
|
||||
|
||||
```json
|
||||
{
|
||||
"your_provider": {
|
||||
"base_url": "https://api.yourprovider.com/v1",
|
||||
"api_key_env": "YOUR_PROVIDER_API_KEY"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
That's it! The provider will be automatically loaded and available.
|
||||
|
||||
### Optional Configuration Fields
|
||||
|
||||
```json
|
||||
{
|
||||
"your_provider": {
|
||||
"base_url": "https://api.yourprovider.com/v1",
|
||||
"api_key_env": "YOUR_PROVIDER_API_KEY",
|
||||
|
||||
// Optional: Override base_url via environment variable
|
||||
"api_base_env": "YOUR_PROVIDER_API_BASE",
|
||||
|
||||
// Optional: Which base class to use (default: "openai_gpt")
|
||||
"base_class": "openai_gpt", // or "openai_like"
|
||||
|
||||
// Optional: Parameter name mappings
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
},
|
||||
|
||||
// Optional: Parameter constraints
|
||||
"constraints": {
|
||||
"temperature_max": 1.0,
|
||||
"temperature_min": 0.0,
|
||||
"temperature_min_with_n_gt_1": 0.3
|
||||
},
|
||||
|
||||
// Optional: Special handling flags
|
||||
"special_handling": {
|
||||
"convert_content_list_to_string": true
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Example: PublicAI
|
||||
|
||||
The first JSON-configured provider:
|
||||
|
||||
```json
|
||||
{
|
||||
"publicai": {
|
||||
"base_url": "https://api.publicai.co/v1",
|
||||
"api_key_env": "PUBLICAI_API_KEY",
|
||||
"api_base_env": "PUBLICAI_API_BASE",
|
||||
"base_class": "openai_gpt",
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
},
|
||||
"special_handling": {
|
||||
"convert_content_list_to_string": true
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
response = litellm.completion(
|
||||
model="publicai/swiss-ai/apertus-8b-instruct",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
)
|
||||
```
|
||||
|
||||
## Benefits
|
||||
|
||||
- **Simple**: 2-5 lines of JSON vs 100+ lines of Python
|
||||
- **Fast**: Add a provider in 5 minutes
|
||||
- **Safe**: No Python code to mess up
|
||||
- **Consistent**: All providers follow the same pattern
|
||||
- **Maintainable**: Centralized configuration
|
||||
|
||||
## When to Use Python Instead
|
||||
|
||||
Use a Python config class if you need:
|
||||
- Custom authentication (OAuth, rotating tokens, etc.)
|
||||
- Complex request/response transformations
|
||||
- Provider-specific streaming logic
|
||||
- Advanced tool calling transformations
|
||||
|
||||
## Implementation Details
|
||||
|
||||
### How It Works
|
||||
|
||||
1. `json_loader.py` loads `providers.json` on import
|
||||
2. `dynamic_config.py` generates config classes on-demand
|
||||
3. Provider resolution checks JSON registry first
|
||||
4. ProviderConfigManager returns JSON-based configs
|
||||
|
||||
### Integration Points
|
||||
|
||||
The JSON system is integrated at:
|
||||
- `litellm/litellm_core_utils/get_llm_provider_logic.py` - Provider resolution
|
||||
- `litellm/utils.py` - ProviderConfigManager
|
||||
- `litellm/constants.py` - openai_compatible_providers list
|
||||
145
litellm/llms/openai_like/dynamic_config.py
Normal file
145
litellm/llms/openai_like/dynamic_config.py
Normal file
|
|
@ -0,0 +1,145 @@
|
|||
"""
|
||||
Dynamic configuration class generator for JSON-based providers.
|
||||
"""
|
||||
|
||||
from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, overload
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
handle_messages_with_content_list_to_str_conversion,
|
||||
)
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
from .json_loader import SimpleProviderConfig
|
||||
|
||||
|
||||
def create_config_class(provider: SimpleProviderConfig):
|
||||
"""Generate config class dynamically from JSON configuration"""
|
||||
|
||||
# Choose base class
|
||||
base_class = (
|
||||
OpenAIGPTConfig if provider.base_class == "openai_gpt" else OpenAILikeChatConfig
|
||||
)
|
||||
|
||||
class JSONProviderConfig(base_class):
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: Literal[True]
|
||||
) -> Coroutine[Any, Any, List[AllMessageValues]]:
|
||||
...
|
||||
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self,
|
||||
messages: List[AllMessageValues],
|
||||
model: str,
|
||||
is_async: Literal[False] = False,
|
||||
) -> List[AllMessageValues]:
|
||||
...
|
||||
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: bool = False
|
||||
) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]:
|
||||
"""Transform messages based on special_handling config"""
|
||||
|
||||
# Handle content list to string conversion if configured
|
||||
if provider.special_handling.get("convert_content_list_to_string"):
|
||||
messages = handle_messages_with_content_list_to_str_conversion(messages)
|
||||
|
||||
if is_async:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=True
|
||||
)
|
||||
else:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=False
|
||||
)
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""Get API base and key from JSON config"""
|
||||
|
||||
# Resolve base URL
|
||||
resolved_base = api_base
|
||||
if not resolved_base and provider.api_base_env:
|
||||
resolved_base = get_secret_str(provider.api_base_env)
|
||||
if not resolved_base:
|
||||
resolved_base = provider.base_url
|
||||
|
||||
# Resolve API key
|
||||
resolved_key = api_key or get_secret_str(provider.api_key_env)
|
||||
|
||||
return resolved_base, resolved_key
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""Build complete URL for the API endpoint"""
|
||||
if not api_base:
|
||||
api_base = provider.base_url
|
||||
|
||||
if not api_base.endswith("/chat/completions"):
|
||||
api_base = f"{api_base}/chat/completions"
|
||||
|
||||
return api_base
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
"""Get supported OpenAI params from base class"""
|
||||
return super().get_supported_openai_params(model=model)
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""Apply parameter mappings and constraints"""
|
||||
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
# Apply supported params
|
||||
for param, value in non_default_params.items():
|
||||
# Check parameter mappings first
|
||||
if param in provider.param_mappings:
|
||||
optional_params[provider.param_mappings[param]] = value
|
||||
elif param in supported_params:
|
||||
optional_params[param] = value
|
||||
|
||||
# Apply temperature constraints if present
|
||||
if "temperature" in optional_params:
|
||||
temp = optional_params["temperature"]
|
||||
constraints = provider.constraints
|
||||
|
||||
# Clamp to max
|
||||
if "temperature_max" in constraints:
|
||||
temp = min(temp, constraints["temperature_max"])
|
||||
|
||||
# Clamp to min
|
||||
if "temperature_min" in constraints:
|
||||
temp = max(temp, constraints["temperature_min"])
|
||||
|
||||
# Special case: temperature_min_with_n_gt_1
|
||||
if "temperature_min_with_n_gt_1" in constraints:
|
||||
n = optional_params.get("n", 1)
|
||||
if n > 1 and temp < constraints["temperature_min_with_n_gt_1"]:
|
||||
temp = constraints["temperature_min_with_n_gt_1"]
|
||||
|
||||
optional_params["temperature"] = temp
|
||||
|
||||
return optional_params
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return provider.slug
|
||||
|
||||
return JSONProviderConfig
|
||||
72
litellm/llms/openai_like/json_loader.py
Normal file
72
litellm/llms/openai_like/json_loader.py
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
"""
|
||||
JSON-based provider configuration loader for OpenAI-compatible providers.
|
||||
"""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Dict, Optional
|
||||
|
||||
|
||||
class SimpleProviderConfig:
|
||||
"""Simple data class for JSON provider config"""
|
||||
|
||||
def __init__(self, slug: str, data: dict):
|
||||
self.slug = slug
|
||||
self.base_url = data["base_url"]
|
||||
self.api_key_env = data["api_key_env"]
|
||||
self.api_base_env = data.get("api_base_env")
|
||||
self.base_class = data.get("base_class", "openai_gpt")
|
||||
self.param_mappings = data.get("param_mappings", {})
|
||||
self.constraints = data.get("constraints", {})
|
||||
self.special_handling = data.get("special_handling", {})
|
||||
|
||||
|
||||
class JSONProviderRegistry:
|
||||
"""Load providers from JSON once on import"""
|
||||
|
||||
_providers: Dict[str, SimpleProviderConfig] = {}
|
||||
_loaded = False
|
||||
|
||||
@classmethod
|
||||
def load(cls):
|
||||
"""Load providers from JSON configuration file"""
|
||||
if cls._loaded:
|
||||
return
|
||||
|
||||
json_path = Path(__file__).parent / "providers.json"
|
||||
|
||||
if not json_path.exists():
|
||||
# No JSON file yet, that's okay
|
||||
cls._loaded = True
|
||||
return
|
||||
|
||||
try:
|
||||
with open(json_path) as f:
|
||||
data = json.load(f)
|
||||
|
||||
for slug, config in data.items():
|
||||
cls._providers[slug] = SimpleProviderConfig(slug, config)
|
||||
|
||||
cls._loaded = True
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to load JSON provider configs: {e}")
|
||||
cls._loaded = True
|
||||
|
||||
@classmethod
|
||||
def get(cls, slug: str) -> Optional[SimpleProviderConfig]:
|
||||
"""Get a provider configuration by slug"""
|
||||
return cls._providers.get(slug)
|
||||
|
||||
@classmethod
|
||||
def exists(cls, slug: str) -> bool:
|
||||
"""Check if a provider is defined via JSON"""
|
||||
return slug in cls._providers
|
||||
|
||||
@classmethod
|
||||
def list_providers(cls) -> list:
|
||||
"""List all registered provider slugs"""
|
||||
return list(cls._providers.keys())
|
||||
|
||||
|
||||
# Load on import
|
||||
JSONProviderRegistry.load()
|
||||
14
litellm/llms/openai_like/providers.json
Normal file
14
litellm/llms/openai_like/providers.json
Normal file
|
|
@ -0,0 +1,14 @@
|
|||
{
|
||||
"publicai": {
|
||||
"base_url": "https://api.publicai.co/v1",
|
||||
"api_key_env": "PUBLICAI_API_KEY",
|
||||
"api_base_env": "PUBLICAI_API_BASE",
|
||||
"base_class": "openai_gpt",
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
},
|
||||
"special_handling": {
|
||||
"convert_content_list_to_string": true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -1,114 +0,0 @@
|
|||
"""
|
||||
Translates from OpenAI's `/v1/chat/completions` to PublicAI's `/v1/chat/completions`
|
||||
"""
|
||||
|
||||
from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, overload
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
handle_messages_with_content_list_to_str_conversion,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
|
||||
class PublicAIChatConfig(OpenAIGPTConfig):
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: Literal[True]
|
||||
) -> Coroutine[Any, Any, List[AllMessageValues]]:
|
||||
...
|
||||
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self,
|
||||
messages: List[AllMessageValues],
|
||||
model: str,
|
||||
is_async: Literal[False] = False,
|
||||
) -> List[AllMessageValues]:
|
||||
...
|
||||
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: bool = False
|
||||
) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]:
|
||||
"""
|
||||
PublicAI does not support content in list format.
|
||||
"""
|
||||
messages = handle_messages_with_content_list_to_str_conversion(messages)
|
||||
if is_async:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=True
|
||||
)
|
||||
else:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=False
|
||||
)
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
api_base = (
|
||||
api_base
|
||||
or get_secret_str("PUBLICAI_API_BASE")
|
||||
or "https://platform.publicai.co/v1"
|
||||
) # type: ignore
|
||||
dynamic_api_key = api_key or get_secret_str("PUBLICAI_API_KEY")
|
||||
return api_base, dynamic_api_key
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
If api_base is not provided, use the default PublicAI /chat/completions endpoint.
|
||||
"""
|
||||
if not api_base:
|
||||
api_base = "https://platform.publicai.co/v1"
|
||||
|
||||
if not api_base.endswith("/chat/completions"):
|
||||
api_base = f"{api_base}/chat/completions"
|
||||
|
||||
return api_base
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
"""
|
||||
Get the supported OpenAI params for PublicAI models
|
||||
|
||||
PublicAI limitations:
|
||||
- functions parameter is not supported (use tools instead)
|
||||
"""
|
||||
excluded_params: List[str] = ["functions"]
|
||||
|
||||
base_openai_params = super().get_supported_openai_params(model=model)
|
||||
final_params: List[str] = []
|
||||
for param in base_openai_params:
|
||||
if param not in excluded_params:
|
||||
final_params.append(param)
|
||||
|
||||
return final_params
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI parameters to PublicAI parameters
|
||||
"""
|
||||
supported_openai_params = self.get_supported_openai_params(model)
|
||||
for param, value in non_default_params.items():
|
||||
if param == "max_completion_tokens":
|
||||
optional_params["max_tokens"] = value
|
||||
elif param in supported_openai_params:
|
||||
optional_params[param] = value
|
||||
|
||||
return optional_params
|
||||
|
||||
|
|
@ -7022,6 +7022,14 @@ class ProviderConfigManager:
|
|||
Returns the provider config for a given provider.
|
||||
"""
|
||||
|
||||
# Check JSON providers FIRST
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
|
||||
if JSONProviderRegistry.exists(provider.value):
|
||||
provider_config = JSONProviderRegistry.get(provider.value)
|
||||
return create_config_class(provider_config)()
|
||||
|
||||
if (
|
||||
provider == LlmProviders.OPENAI
|
||||
and litellm.openaiOSeriesConfig.is_model_o_series_model(model=model)
|
||||
|
|
|
|||
0
tests/test_litellm/llms/openai_like/__init__.py
Normal file
0
tests/test_litellm/llms/openai_like/__init__.py
Normal file
311
tests/test_litellm/llms/openai_like/test_json_providers.py
Normal file
311
tests/test_litellm/llms/openai_like/test_json_providers.py
Normal file
|
|
@ -0,0 +1,311 @@
|
|||
"""
|
||||
Tests for JSON-based provider configuration system.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
try:
|
||||
import pytest
|
||||
except ImportError:
|
||||
# pytest not available, will run as standalone script
|
||||
pytest = None
|
||||
|
||||
# Add workspace to path
|
||||
workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))
|
||||
sys.path.insert(0, workspace_path)
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
class TestJSONProviderLoader:
|
||||
"""Test JSON provider loading and configuration"""
|
||||
|
||||
def test_load_json_providers(self):
|
||||
"""Test that JSON providers load correctly"""
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
# Verify publicai is loaded
|
||||
assert JSONProviderRegistry.exists("publicai")
|
||||
|
||||
# Get publicai config
|
||||
publicai = JSONProviderRegistry.get("publicai")
|
||||
assert publicai is not None
|
||||
assert publicai.base_url == "https://api.publicai.co/v1"
|
||||
assert publicai.api_key_env == "PUBLICAI_API_KEY"
|
||||
assert publicai.api_base_env == "PUBLICAI_API_BASE"
|
||||
assert publicai.param_mappings.get("max_completion_tokens") == "max_tokens"
|
||||
|
||||
def test_dynamic_config_generation(self):
|
||||
"""Test dynamic config class creation"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("publicai")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
# Test API info resolution
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
|
||||
assert api_base == "https://api.publicai.co/v1"
|
||||
|
||||
# Test with custom base
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(
|
||||
"https://custom.api.com", "test-key"
|
||||
)
|
||||
assert api_base == "https://custom.api.com"
|
||||
assert api_key == "test-key"
|
||||
|
||||
def test_parameter_mapping(self):
|
||||
"""Test parameter mapping works"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("publicai")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
# Test parameter mapping
|
||||
optional_params = {}
|
||||
non_default_params = {"max_completion_tokens": 100, "temperature": 0.7}
|
||||
result = config.map_openai_params(
|
||||
non_default_params, optional_params, "gpt-4", False
|
||||
)
|
||||
|
||||
# max_completion_tokens should be mapped to max_tokens
|
||||
assert "max_tokens" in result
|
||||
assert result["max_tokens"] == 100
|
||||
assert "max_completion_tokens" not in result
|
||||
|
||||
# temperature should be passed through
|
||||
assert result["temperature"] == 0.7
|
||||
|
||||
def test_supported_params(self):
|
||||
"""Test that config returns supported params"""
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
|
||||
provider = JSONProviderRegistry.get("publicai")
|
||||
config_class = create_config_class(provider)
|
||||
config = config_class()
|
||||
|
||||
# Get supported params
|
||||
supported = config.get_supported_openai_params("gpt-4")
|
||||
|
||||
# Should have standard OpenAI params
|
||||
assert isinstance(supported, list)
|
||||
assert len(supported) > 0
|
||||
|
||||
def test_provider_resolution(self):
|
||||
"""Test that provider resolution finds JSON providers"""
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import (
|
||||
get_llm_provider,
|
||||
)
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="publicai/gpt-4",
|
||||
custom_llm_provider=None,
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
assert model == "gpt-4"
|
||||
assert provider == "publicai"
|
||||
assert api_base == "https://api.publicai.co/v1"
|
||||
|
||||
def test_provider_config_manager(self):
|
||||
"""Test that ProviderConfigManager returns JSON-based configs"""
|
||||
from litellm import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
config = ProviderConfigManager.get_provider_chat_config(
|
||||
model="gpt-4", provider=LlmProviders.PUBLICAI
|
||||
)
|
||||
|
||||
assert config is not None
|
||||
assert config.custom_llm_provider == "publicai"
|
||||
|
||||
|
||||
class TestPublicAIIntegration:
|
||||
"""Integration tests for PublicAI provider"""
|
||||
|
||||
def test_publicai_completion_basic(self):
|
||||
"""Test basic completion call to PublicAI"""
|
||||
# Set API key from the one provided
|
||||
os.environ["PUBLICAI_API_KEY"] = (
|
||||
"zpka_9ea399e9e81b4ece8af0fe88d2561c4f_4e4e9dec"
|
||||
)
|
||||
|
||||
try:
|
||||
response = litellm.completion(
|
||||
model="publicai/swiss-ai/apertus-8b-instruct",
|
||||
messages=[{"role": "user", "content": "Say 'test successful' and nothing else"}],
|
||||
max_tokens=10,
|
||||
)
|
||||
|
||||
# Verify response structure
|
||||
assert response is not None
|
||||
assert hasattr(response, "choices")
|
||||
assert len(response.choices) > 0
|
||||
assert hasattr(response.choices[0], "message")
|
||||
assert hasattr(response.choices[0].message, "content")
|
||||
assert response.choices[0].message.content is not None
|
||||
|
||||
# Check that we got a response
|
||||
content = response.choices[0].message.content.lower()
|
||||
assert len(content) > 0
|
||||
|
||||
print(f"✓ PublicAI completion successful: {response.choices[0].message.content}")
|
||||
|
||||
except Exception as e:
|
||||
if pytest:
|
||||
pytest.fail(f"PublicAI completion failed: {str(e)}")
|
||||
else:
|
||||
raise
|
||||
|
||||
def test_publicai_completion_with_streaming(self):
|
||||
"""Test streaming completion with PublicAI"""
|
||||
os.environ["PUBLICAI_API_KEY"] = (
|
||||
"zpka_9ea399e9e81b4ece8af0fe88d2561c4f_4e4e9dec"
|
||||
)
|
||||
|
||||
try:
|
||||
response = litellm.completion(
|
||||
model="publicai/swiss-ai/apertus-8b-instruct",
|
||||
messages=[{"role": "user", "content": "Count to 3"}],
|
||||
max_tokens=20,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
# Collect chunks
|
||||
chunks = []
|
||||
for chunk in response:
|
||||
assert chunk is not None
|
||||
if hasattr(chunk.choices[0], "delta") and hasattr(
|
||||
chunk.choices[0].delta, "content"
|
||||
):
|
||||
if chunk.choices[0].delta.content:
|
||||
chunks.append(chunk.choices[0].delta.content)
|
||||
|
||||
# Verify we got chunks
|
||||
assert len(chunks) > 0
|
||||
full_response = "".join(chunks)
|
||||
assert len(full_response) > 0
|
||||
|
||||
print(f"✓ PublicAI streaming successful: {full_response}")
|
||||
|
||||
except Exception as e:
|
||||
if pytest:
|
||||
pytest.fail(f"PublicAI streaming failed: {str(e)}")
|
||||
else:
|
||||
raise
|
||||
|
||||
def test_publicai_parameter_mapping(self):
|
||||
"""Test that max_completion_tokens is mapped to max_tokens"""
|
||||
os.environ["PUBLICAI_API_KEY"] = (
|
||||
"zpka_9ea399e9e81b4ece8af0fe88d2561c4f_4e4e9dec"
|
||||
)
|
||||
|
||||
try:
|
||||
# Use max_completion_tokens (OpenAI's newer parameter)
|
||||
response = litellm.completion(
|
||||
model="publicai/swiss-ai/apertus-8b-instruct",
|
||||
messages=[{"role": "user", "content": "Hi"}],
|
||||
max_completion_tokens=5, # This should be mapped to max_tokens
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert len(response.choices) > 0
|
||||
|
||||
print("✓ Parameter mapping successful")
|
||||
|
||||
except Exception as e:
|
||||
if pytest:
|
||||
pytest.fail(f"Parameter mapping test failed: {str(e)}")
|
||||
else:
|
||||
raise
|
||||
|
||||
def test_publicai_content_list_conversion(self):
|
||||
"""Test that content list format is converted to string"""
|
||||
os.environ["PUBLICAI_API_KEY"] = (
|
||||
"zpka_9ea399e9e81b4ece8af0fe88d2561c4f_4e4e9dec"
|
||||
)
|
||||
|
||||
try:
|
||||
# Send message with content as list (should be converted to string)
|
||||
response = litellm.completion(
|
||||
model="publicai/swiss-ai/apertus-8b-instruct",
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Say hello"}
|
||||
]
|
||||
}
|
||||
],
|
||||
max_tokens=10,
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert len(response.choices) > 0
|
||||
|
||||
print("✓ Content list conversion successful")
|
||||
|
||||
except Exception as e:
|
||||
if pytest:
|
||||
pytest.fail(f"Content list conversion test failed: {str(e)}")
|
||||
else:
|
||||
raise
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run basic tests
|
||||
print("Testing JSON Provider System...")
|
||||
|
||||
test_loader = TestJSONProviderLoader()
|
||||
print("\n1. Testing JSON provider loading...")
|
||||
test_loader.test_load_json_providers()
|
||||
print(" ✓ JSON providers loaded")
|
||||
|
||||
print("\n2. Testing dynamic config generation...")
|
||||
test_loader.test_dynamic_config_generation()
|
||||
print(" ✓ Dynamic config works")
|
||||
|
||||
print("\n3. Testing parameter mapping...")
|
||||
test_loader.test_parameter_mapping()
|
||||
print(" ✓ Parameter mapping works")
|
||||
|
||||
print("\n4. Testing excluded params...")
|
||||
test_loader.test_excluded_params()
|
||||
print(" ✓ Excluded params work")
|
||||
|
||||
print("\n5. Testing provider resolution...")
|
||||
test_loader.test_provider_resolution()
|
||||
print(" ✓ Provider resolution works")
|
||||
|
||||
print("\n6. Testing provider config manager...")
|
||||
test_loader.test_provider_config_manager()
|
||||
print(" ✓ Config manager works")
|
||||
|
||||
print("\n" + "="*50)
|
||||
print("PublicAI Integration Tests...")
|
||||
print("="*50)
|
||||
|
||||
test_integration = TestPublicAIIntegration()
|
||||
|
||||
print("\n7. Testing basic completion...")
|
||||
test_integration.test_publicai_completion_basic()
|
||||
|
||||
print("\n8. Testing streaming...")
|
||||
test_integration.test_publicai_completion_with_streaming()
|
||||
|
||||
print("\n9. Testing parameter mapping...")
|
||||
test_integration.test_publicai_parameter_mapping()
|
||||
|
||||
print("\n10. Testing content list conversion...")
|
||||
test_integration.test_publicai_content_list_conversion()
|
||||
|
||||
print("\n" + "="*50)
|
||||
print("✓ All tests passed!")
|
||||
print("="*50)
|
||||
Loading…
Add table
Reference in a new issue