make generic api OSS + support multiple generic API's (#17152)

* feat(generic_api_callback.py): make generic api OSS + support multiple generic API's

Enables https://github.com/BerriAI/litellm/pull/17094#discussion_r2562832967

* feat(callback_utils.py): support custom generic api callbacks

* feat(generic_api_callback.py): support specifying which event types to run the generic api for

* fix(litellm_logging.py): log system prompt for anthropic messages

* feat(generic_api_callback.py): support generic api compatible api's - e.g. rubrik agent cloud

* docs(sidebars.js): document new OSS generic api

* docs(generic_api.md): document new OSS Generic API

* docs(custom_webhook_api.md): document custom webhook api integration tutorial

* docs(custom_webhook_api.md): cleanup

* docs(custom_webhook_api.md): document what get's logged to custom webhook api

* Refactor: Pass callback config to GenericAPILogger

Co-authored-by: krrishdholakia <krrishdholakia@gmail.com>

* Fix: Handle empty messages list in logging payload

Co-authored-by: krrishdholakia <krrishdholakia@gmail.com>

* Checkpoint before follow-up message

Co-authored-by: krrishdholakia <krrishdholakia@gmail.com>

* feat: Cache GenericAPILogger instances to improve performance

Co-authored-by: krrishdholakia <krrishdholakia@gmail.com>

---------

Co-authored-by: Cursor Agent <cursoragent@cursor.com>
This commit is contained in:
Krish Dholakia
2025-11-26 18:38:38 -08:00
committed by GitHub
co-authored by Cursor Agent
parent 48eb34a875
commit 1cb5fcddba
13 changed files with 560 additions and 79 deletions
@@ -0,0 +1,106 @@
# Contribute Custom Webhook API
If your API just needs a Webhook event from LiteLLM, here's how to add a 'native' integration for it on LiteLLM:
1. Clone the repo and open the `generic_api_compatible_callbacks.json`
```bash
git clone https://github.com/BerriAI/litellm.git
cd litellm
open .
```
2. Add your API to the `generic_api_compatible_callbacks.json`
Example:
```json
{
"rubrik": {
"event_types": ["llm_api_success"],
"endpoint": "{{environment_variables.RUBRIK_WEBHOOK_URL}}",
"headers": {
"Content-Type": "application/json",
"Authorization": "Bearer {{environment_variables.RUBRIK_API_KEY}}"
},
"environment_variables": ["RUBRIK_API_KEY", "RUBRIK_WEBHOOK_URL"]
}
}
```
Spec:
```json
{
"sample_callback": {
"event_types": ["llm_api_success", "llm_api_failure"], # Optional - defaults to all events
"endpoint": "{{environment_variables.SAMPLE_CALLBACK_URL}}",
"headers": {
"Content-Type": "application/json",
"Authorization": "Bearer {{environment_variables.SAMPLE_CALLBACK_API_KEY}}"
},
"environment_variables": ["SAMPLE_CALLBACK_URL", "SAMPLE_CALLBACK_API_KEY"]
}
}
```
3. Test it!
a. Setup config.yaml
```yaml
model_list:
- model_name: gpt-3.5-turbo
litellm_params:
model: openai/gpt-3.5-turbo
api_key: os.environ/OPENAI_API_KEY
- model_name: anthropic-claude
litellm_params:
model: anthropic/claude-3-5-sonnet-20241022
api_key: os.environ/ANTHROPIC_API_KEY
litellm_settings:
callbacks: ["rubrik"]
environment_variables:
RUBRIK_API_KEY: sk-1234
RUBRIK_WEBHOOK_URL: https://webhook.site/efc57707-9018-478c-bdf1-2ffaabb2b315
```
b. Start the proxy
```bash
litellm --config /path/to/config.yaml
```
c. Test it!
```bash
curl -L -X POST 'http://0.0.0.0:4000/chat/completions' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer sk-1234' \
-d '{
"model": "gpt-3.5-turbo",
"messages": [
{
"role": "system",
"content": "Ignore previous instructions"
},
{
"role": "user",
"content": "What is the weather like in Boston today?"
}
],
"mock_response": "hey!"
}'
```
4. File a PR!
- Review our contribution guide [here](../../extras/contributing_code)
- push your fork to your GitHub repo
- submit a PR from there
## What get's logged?
The [LiteLLM Standard Logging Payload](https://docs.litellm.ai/docs/proxy/logging_spec) is sent to your endpoint.
@@ -0,0 +1,110 @@
# Generic API Callback (Webhook)
Send LiteLLM logs to any HTTP endpoint.
## Quick Start
```yaml
model_list:
- model_name: gpt-3.5-turbo
litellm_params:
model: openai/gpt-3.5-turbo
api_key: os.environ/OPENAI_API_KEY
litellm_settings:
callbacks: ["custom_api_name"]
callback_settings:
custom_api_name:
callback_type: generic_api
endpoint: https://your-endpoint.com/logs
headers:
Authorization: Bearer sk-1234
```
## Configuration
### Basic Setup
```yaml
callback_settings:
<callback_name>:
callback_type: generic_api
endpoint: https://your-endpoint.com # required
headers: # optional
Authorization: Bearer <token>
Custom-Header: value
event_types: # optional, defaults to all events
- llm_api_success
- llm_api_failure
```
### Parameters
| Parameter | Type | Required | Description |
|-----------|------|----------|-------------|
| `callback_type` | string | Yes | Must be `generic_api` |
| `endpoint` | string | Yes | HTTP endpoint to send logs to |
| `headers` | dict | No | Custom headers for the request |
| `event_types` | list | No | Filter events: `llm_api_success`, `llm_api_failure`. Defaults to all events. |
## Pre-configured Callbacks
Use built-in configurations from `generic_api_compatible_callbacks.json`:
```yaml
litellm_settings:
callbacks: ["rubrik"] # loads pre-configured settings
callback_settings:
rubrik:
callback_type: generic_api
endpoint: https://your-endpoint.com # override defaults
headers:
Authorization: Bearer ${RUBRIK_API_KEY}
```
## Payload Format
Logs are sent as `StandardLoggingPayload` [objects](https://docs.litellm.ai/docs/proxy/logging_spec) in JSON format:
```json
[
{
"id": "chatcmpl-123",
"call_type": "litellm.completion",
"model": "gpt-3.5-turbo",
"messages": [...],
"response": {...},
"usage": {...},
"cost": 0.0001,
"startTime": "2024-01-01T00:00:00",
"endTime": "2024-01-01T00:00:01",
"metadata": {...}
}
]
```
## Environment Variables
Set via environment variables instead of config:
```bash
export GENERIC_LOGGER_ENDPOINT=https://your-endpoint.com
export GENERIC_LOGGER_HEADERS="Authorization=Bearer token,Custom-Header=value"
```
## Batch Settings
Control batching behavior (inherits from `CustomBatchLogger`):
```yaml
callback_settings:
my_api:
callback_type: generic_api
endpoint: https://your-endpoint.com
batch_size: 100 # default: 100
flush_interval: 60 # seconds, default: 60
```
+10
View File
@@ -20,6 +20,16 @@ const sidebars = {
type: "category",
label: "Observability",
items: [
{
type: "category",
label: "Contributing to Integrations",
items: [
{
type: "autogenerated",
dirName: "contribute_integration"
}
]
},
{
type: "autogenerated",
dirName: "observability"
+1
View File
@@ -177,6 +177,7 @@ _known_custom_logger_compatible_callbacks: List = list(
callbacks: List[
Union[Callable, _custom_logger_compatible_callbacks_literal, CustomLogger]
] = []
callback_settings: Dict[str, Dict[str, Any]] = {}
initialized_langfuse_clients: int = 0
langfuse_default_tags: Optional[List[str]] = None
langsmith_batch_size: Optional[int] = None
@@ -7,13 +7,15 @@ Callback to log events to a Generic API Endpoint
"""
import asyncio
import json
import os
import re
import traceback
from litellm._uuid import uuid
from typing import Dict, List, Optional, Union
from typing import Dict, List, Literal, Optional, Union
import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import (
@@ -22,12 +24,83 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.types.utils import StandardLoggingPayload
API_EVENT_TYPES = Literal["llm_api_success", "llm_api_failure"]
def load_compatible_callbacks() -> Dict:
"""
Load the generic_api_compatible_callbacks.json file
Returns:
Dict: Dictionary of compatible callbacks configuration
"""
try:
json_path = os.path.join(
os.path.dirname(__file__), "generic_api_compatible_callbacks.json"
)
with open(json_path, "r") as f:
return json.load(f)
except Exception as e:
verbose_logger.warning(
f"Error loading generic_api_compatible_callbacks.json: {str(e)}"
)
return {}
def is_callback_compatible(callback_name: str) -> bool:
"""
Check if a callback_name exists in the compatible callbacks list
Args:
callback_name: Name of the callback to check
Returns:
bool: True if callback_name exists in the compatible callbacks, False otherwise
"""
compatible_callbacks = load_compatible_callbacks()
return callback_name in compatible_callbacks
def get_callback_config(callback_name: str) -> Optional[Dict]:
"""
Get the configuration for a specific callback
Args:
callback_name: Name of the callback to get config for
Returns:
Optional[Dict]: Configuration dict for the callback, or None if not found
"""
compatible_callbacks = load_compatible_callbacks()
return compatible_callbacks.get(callback_name)
def substitute_env_variables(value: str) -> str:
"""
Replace {{environment_variables.VAR_NAME}} patterns with actual environment variable values
Args:
value: String that may contain {{environment_variables.VAR_NAME}} patterns
Returns:
str: String with environment variables substituted
"""
pattern = r"\{\{environment_variables\.([A-Z_]+)\}\}"
def replace_env_var(match):
env_var_name = match.group(1)
return os.getenv(env_var_name, "")
return re.sub(pattern, replace_env_var, value)
class GenericAPILogger(CustomBatchLogger):
def __init__(
self,
endpoint: Optional[str] = None,
headers: Optional[dict] = None,
event_types: Optional[List[API_EVENT_TYPES]] = None,
callback_name: Optional[str] = None,
**kwargs,
):
"""
@@ -36,7 +109,37 @@ class GenericAPILogger(CustomBatchLogger):
Args:
endpoint: Optional[str] = None,
headers: Optional[dict] = None,
event_types: Optional[List[API_EVENT_TYPES]] = None,
callback_name: Optional[str] = None - If provided, loads config from generic_api_compatible_callbacks.json
"""
#########################################################
# Check if callback_name is provided and load config
#########################################################
if callback_name:
if is_callback_compatible(callback_name):
verbose_logger.debug(
f"Loading configuration for callback: {callback_name}"
)
callback_config = get_callback_config(callback_name)
# Use config from JSON if not explicitly provided
if callback_config:
if endpoint is None and "endpoint" in callback_config:
endpoint = substitute_env_variables(callback_config["endpoint"])
if "headers" in callback_config:
headers = headers or {}
for key, value in callback_config["headers"].items():
if key not in headers:
headers[key] = substitute_env_variables(value)
if event_types is None and "event_types" in callback_config:
event_types = callback_config["event_types"]
else:
verbose_logger.warning(
f"callback_name '{callback_name}' not found in generic_api_compatible_callbacks.json"
)
#########################################################
# Init httpx client
#########################################################
@@ -51,8 +154,10 @@ class GenericAPILogger(CustomBatchLogger):
self.headers: Dict = self._get_headers(headers)
self.endpoint: str = endpoint
self.event_types: Optional[List[API_EVENT_TYPES]] = event_types
self.callback_name: Optional[str] = callback_name
verbose_logger.debug(
f"in init GenericAPILogger, endpoint {self.endpoint}, headers {self.headers}"
f"in init GenericAPILogger, callback_name: {self.callback_name}, endpoint {self.endpoint}, headers {self.headers}, event_types: {self.event_types}"
)
#########################################################
@@ -114,9 +219,9 @@ class GenericAPILogger(CustomBatchLogger):
Raises:
Raises a NON Blocking verbose_logger.exception if an error occurs
"""
from litellm.proxy.utils import _premium_user_check
_premium_user_check()
if self.event_types is not None and "llm_api_success" not in self.event_types:
return
try:
verbose_logger.debug(
@@ -153,9 +258,8 @@ class GenericAPILogger(CustomBatchLogger):
- Creates a StandardLoggingPayload
- Adds to batch queue
"""
from litellm.proxy.utils import _premium_user_check
_premium_user_check()
if self.event_types is not None and "llm_api_failure" not in self.event_types:
return
try:
verbose_logger.debug(
@@ -0,0 +1,20 @@
{
"sample_callback": {
"event_types": ["llm_api_success", "llm_api_failure"],
"endpoint": "{{environment_variables.SAMPLE_CALLBACK_URL}}",
"headers": {
"Content-Type": "application/json",
"Authorization": "Bearer {{environment_variables.SAMPLE_CALLBACK_API_KEY}}"
},
"environment_variables": ["SAMPLE_CALLBACK_URL", "SAMPLE_CALLBACK_API_KEY"]
},
"rubrik": {
"event_types": ["llm_api_success"],
"endpoint": "{{environment_variables.RUBRIK_WEBHOOK_URL}}",
"headers": {
"Content-Type": "application/json",
"Authorization": "Bearer {{environment_variables.RUBRIK_API_KEY}}"
},
"environment_variables": ["RUBRIK_API_KEY", "RUBRIK_WEBHOOK_URL"]
}
}
@@ -95,9 +95,6 @@ class CustomLoggerRegistry:
}
try:
from litellm_enterprise.enterprise_callbacks.generic_api_callback import (
GenericAPILogger,
)
from litellm_enterprise.enterprise_callbacks.pagerduty.pagerduty import (
PagerDutyAlerting,
)
@@ -108,6 +105,10 @@ class CustomLoggerRegistry:
SMTPEmailLogger,
)
from litellm.integrations.generic_api.generic_api_callback import (
GenericAPILogger,
)
enterprise_loggers = {
"pagerduty": PagerDutyAlerting,
"generic_api": GenericAPILogger,
+41 -8
View File
@@ -165,9 +165,6 @@ try:
from litellm_enterprise.enterprise_callbacks.callback_controls import (
EnterpriseCallbackControls,
)
from litellm_enterprise.enterprise_callbacks.generic_api_callback import (
GenericAPILogger,
)
from litellm_enterprise.enterprise_callbacks.pagerduty.pagerduty import (
PagerDutyAlerting,
)
@@ -181,6 +178,8 @@ try:
StandardLoggingPayloadSetup as EnterpriseStandardLoggingPayloadSetup,
)
from litellm.integrations.generic_api.generic_api_callback import GenericAPILogger
EnterpriseStandardLoggingPayloadSetupVAR: Optional[
Type[EnterpriseStandardLoggingPayloadSetup]
] = EnterpriseStandardLoggingPayloadSetup
@@ -315,6 +314,7 @@ class Logging(LiteLLMLoggingBaseClass):
for m in messages:
new_messages.append({"role": "user", "content": m})
messages = new_messages
self.model = model
self.messages = copy.deepcopy(messages)
self.stream = stream
@@ -4106,10 +4106,8 @@ def _get_custom_logger_settings_from_proxy_server(callback_name: str) -> Dict:
otel:
message_logging: False
"""
from litellm.proxy.proxy_server import callback_settings
if callback_settings:
return dict(callback_settings.get(callback_name, {}))
if litellm.callback_settings:
return dict(litellm.callback_settings.get(callback_name, {}))
return {}
@@ -4186,6 +4184,39 @@ class StandardLoggingPayloadSetup:
return start_time_float, end_time_float, completion_start_time_float
@staticmethod
def append_system_prompt_messages(
kwargs: Optional[Dict] = None, messages: Optional[Any] = None
):
"""
Append system prompt messages to the messages
"""
if kwargs is not None:
if kwargs.get("system") is not None and isinstance(
kwargs.get("system"), str
):
if messages is None:
return [{"role": "system", "content": kwargs.get("system")}]
elif isinstance(messages, list):
if len(messages) == 0:
return [{"role": "system", "content": kwargs.get("system")}]
# check for duplicates
if messages[0].get("role") == "system" and messages[0].get(
"content"
) == kwargs.get("system"):
return messages
messages = [
{"role": "system", "content": kwargs.get("system")}
] + messages
elif isinstance(messages, str):
messages = [
{"role": "system", "content": kwargs.get("system")},
{"role": "user", "content": messages},
]
return messages
return messages
@staticmethod
def get_standard_logging_metadata(
metadata: Optional[Dict[str, Any]],
@@ -4908,7 +4939,9 @@ def get_standard_logging_object_payload(
model_group=_model_group,
model_id=_model_id,
requester_ip_address=clean_metadata.get("requester_ip_address", None),
messages=kwargs.get("messages"),
messages=StandardLoggingPayloadSetup.append_system_prompt_messages(
kwargs=kwargs, messages=kwargs.get("messages")
),
response=final_response_obj,
model_parameters=ModelParamHelper.get_standard_logging_model_parameters(
kwargs.get("optional_params", None) or {}
@@ -1,9 +1,10 @@
from typing import TYPE_CHECKING, Callable, List, Optional, Set, Type, Union
from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Set, Type, Union
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.additional_logging_utils import AdditionalLoggingUtils
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.generic_api.generic_api_callback import GenericAPILogger
from litellm.types.utils import CallbacksByType
if TYPE_CHECKING:
@@ -11,6 +12,8 @@ if TYPE_CHECKING:
else:
_custom_logger_compatible_callbacks_literal = str
_generic_api_logger_cache: Dict[str, GenericAPILogger] = {}
class LoggingCallbackManager:
"""
@@ -138,6 +141,57 @@ class LoggingCallbackManager:
return False
return True
@staticmethod
def _add_custom_callback_generic_api_str(
callback: str,
) -> Union[GenericAPILogger, str]:
"""
litellm_settings:
success_callback: ["custom_callback_name"]
callback_settings:
custom_callback_name:
callback_type: generic_api
endpoint: https://webhook-test.com/30343bc33591bc5e6dc44217ceae3e0a
headers:
Authorization: Bearer sk-1234
"""
callback_config = litellm.callback_settings.get(callback)
if not isinstance(callback_config, dict):
return callback
if callback_config.get("callback_type") != "generic_api":
return callback
endpoint = callback_config.get("endpoint")
headers = callback_config.get("headers")
event_types = callback_config.get("event_types")
if endpoint is None or headers is None:
verbose_logger.warning(
"generic_api callback '%s' is missing endpoint or headers, skipping.",
callback,
)
return callback
cached_logger = _generic_api_logger_cache.get(callback)
if (
isinstance(cached_logger, GenericAPILogger)
and cached_logger.endpoint == endpoint
and cached_logger.headers == headers
and cached_logger.event_types == event_types
):
return cached_logger
new_logger = GenericAPILogger(
endpoint=endpoint,
headers=headers,
event_types=event_types,
)
_generic_api_logger_cache[callback] = new_logger
return new_logger
def _safe_add_callback_to_list(
self,
callback: Union[CustomLogger, Callable, str],
@@ -152,15 +206,24 @@ class LoggingCallbackManager:
if not self._check_callback_list_size(parent_list):
return
# Check if the callback is a custom callback
if isinstance(callback, str):
callback = LoggingCallbackManager._add_custom_callback_generic_api_str(
callback
)
if isinstance(callback, str):
self._add_string_callback_to_list(
callback=callback, parent_list=parent_list
)
elif isinstance(callback, CustomLogger):
self._add_custom_logger_to_list(
custom_logger=callback,
parent_list=parent_list,
)
elif callable(callback):
self._add_callback_function_to_list(
callback=callback, parent_list=parent_list
@@ -348,7 +411,6 @@ class LoggingCallbackManager:
elif callable(callback):
return getattr(callback, "__name__", str(callback))
return str(callback)
def get_active_custom_logger_for_callback_name(
self,
@@ -362,12 +424,16 @@ class LoggingCallbackManager:
)
# get the custom logger class type
custom_logger_class_type = CustomLoggerRegistry.get_class_type_for_custom_logger_name(callback_name)
custom_logger_class_type = (
CustomLoggerRegistry.get_class_type_for_custom_logger_name(callback_name)
)
# get the active custom logger
custom_logger = self.get_custom_loggers_for_type(custom_logger_class_type)
if len(custom_logger) == 0:
raise ValueError(f"No active custom logger found for callback name: {callback_name}")
raise ValueError(
f"No active custom logger found for callback name: {callback_name}"
)
return custom_logger[0]
+11 -13
View File
@@ -3,19 +3,17 @@ model_list:
litellm_params:
model: openai/gpt-3.5-turbo
api_key: os.environ/OPENAI_API_KEY
guardrails:
- guardrail_name: model-armor-shield
- model_name: anthropic-claude
litellm_params:
guardrail: model_armor
mode: "post_call" # Run on both input and output
template_id: "test-prompt-template" # Required: Your Model Armor template ID
project_id: "test-vector-store-db" # Your GCP project ID
location: "us" # GCP location (default: us-central1)
mask_request_content: true # Enable request content masking
mask_response_content: true # Enable response content masking
fail_on_error: true # Fail request if Model Armor errors (default: true)
default_on: true # Run by default for all requests
model: anthropic/claude-3-5-sonnet-20241022
api_key: os.environ/ANTHROPIC_API_KEY
litellm_settings:
callbacks: ["arize_phoenix"]
callbacks: ["rubrik"]
callback_settings:
rubrik:
callback_type: generic_api
endpoint: https://webhook.site/efc57707-9018-478c-bdf1-2ffaabb2b315
headers:
Authorization: Bearer sk-1234
@@ -24,6 +24,10 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915
litellm_settings: dict,
callback_specific_params: dict = {},
):
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.logging_callback_manager import (
LoggingCallbackManager,
)
from litellm.proxy.proxy_server import prisma_client
verbose_proxy_logger.debug(
@@ -32,6 +36,11 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915
if isinstance(value, list):
imported_list: List[Any] = []
for callback in value: # ["presidio", <my-custom-callback>]
# check if callback is a custom logger compatible callback
if isinstance(callback, str):
callback = LoggingCallbackManager._add_custom_callback_generic_api_str(
callback
)
if (
isinstance(callback, str)
and callback in litellm._known_custom_logger_compatible_callbacks
@@ -259,6 +268,8 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915
**azure_content_safety_params,
)
imported_list.append(azure_content_safety_obj)
elif isinstance(callback, CustomLogger):
imported_list.append(callback)
else:
verbose_proxy_logger.debug(
f"{blue_color_code} attempting to import custom calback={callback} {reset_color_code}"
+10 -12
View File
@@ -286,9 +286,7 @@ from litellm.proxy.management_endpoints.customer_endpoints import (
from litellm.proxy.management_endpoints.internal_user_endpoints import (
router as internal_user_router,
)
from litellm.proxy.management_endpoints.internal_user_endpoints import (
user_update,
)
from litellm.proxy.management_endpoints.internal_user_endpoints import user_update
from litellm.proxy.management_endpoints.key_management_endpoints import (
delete_verification_tokens,
duration_in_seconds,
@@ -342,9 +340,7 @@ from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router
from litellm.proxy.openai_files_endpoints.files_endpoints import (
router as openai_files_router,
)
from litellm.proxy.openai_files_endpoints.files_endpoints import (
set_files_config,
)
from litellm.proxy.openai_files_endpoints.files_endpoints import set_files_config
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
passthrough_endpoint_router,
)
@@ -436,9 +432,7 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
LiteLLM_UpperboundKeyGenerateParams,
)
from litellm.types.realtime import RealtimeQueryParams
from litellm.types.router import (
DeploymentTypedDict,
)
from litellm.types.router import DeploymentTypedDict
from litellm.types.router import ModelInfo as RouterModelInfo
from litellm.types.router import (
RouterGeneralSettings,
@@ -1084,7 +1078,6 @@ llm_router: Optional[Router] = None
llm_model_list: Optional[list] = None
general_settings: dict = {}
config_passthrough_endpoints: Optional[List[Dict[str, Any]]] = None
callback_settings: dict = {}
log_file = "api_log.json"
worker_config = None
master_key: Optional[str] = None
@@ -1241,7 +1234,10 @@ def cost_tracking():
global prisma_client
if prisma_client is not None:
litellm.logging_callback_manager.add_litellm_callback(_ProxyDBLogger())
litellm.logging_callback_manager.add_litellm_async_success_callback(_ProxyDBLogger())
litellm.logging_callback_manager.add_litellm_async_success_callback(
_ProxyDBLogger()
)
async def update_cache( # noqa: PLR0915
token: Optional[str],
@@ -2074,7 +2070,7 @@ class ProxyConfig:
"""
Load config values into proxy global state
"""
global master_key, user_config_file_path, otel_logging, user_custom_auth, user_custom_auth_path, user_custom_key_generate, user_custom_sso, user_custom_ui_sso_sign_in_handler, use_background_health_checks, use_shared_health_check, health_check_interval, use_queue, proxy_budget_rescheduler_max_time, proxy_budget_rescheduler_min_time, ui_access_mode, litellm_master_key_hash, proxy_batch_write_at, disable_spend_logs, prompt_injection_detection_obj, redis_usage_cache, store_model_in_db, premium_user, open_telemetry_logger, health_check_details, callback_settings, proxy_batch_polling_interval, config_passthrough_endpoints
global master_key, user_config_file_path, otel_logging, user_custom_auth, user_custom_auth_path, user_custom_key_generate, user_custom_sso, user_custom_ui_sso_sign_in_handler, use_background_health_checks, use_shared_health_check, health_check_interval, use_queue, proxy_budget_rescheduler_max_time, proxy_budget_rescheduler_min_time, ui_access_mode, litellm_master_key_hash, proxy_batch_write_at, disable_spend_logs, prompt_injection_detection_obj, redis_usage_cache, store_model_in_db, premium_user, open_telemetry_logger, health_check_details, proxy_batch_polling_interval, config_passthrough_endpoints
config: dict = await self.get_config(config_file_path=config_file_path)
@@ -2082,6 +2078,8 @@ class ProxyConfig:
## Callback settings
callback_settings = config.get("callback_settings", {})
if callback_settings:
litellm.callback_settings = callback_settings
## LITELLM MODULE SETTINGS (e.g. litellm.drop_params=True,..)
litellm_settings = config.get("litellm_settings", None)
@@ -28,8 +28,7 @@ from litellm.types.utils import (
)
verbose_logger.setLevel(logging.DEBUG)
from litellm_enterprise.enterprise_callbacks.generic_api_callback import GenericAPILogger
from litellm.integrations.generic_api.generic_api_callback import GenericAPILogger
@pytest.mark.asyncio
@@ -52,9 +51,7 @@ async def test_generic_api_callback():
# Initialize the GenericAPILogger and set the mock
generic_logger = GenericAPILogger(
endpoint=test_endpoint,
headers=test_headers,
flush_interval=1
endpoint=test_endpoint, headers=test_headers, flush_interval=1
)
generic_logger.async_httpx_client.post = mock_post
litellm.callbacks = [generic_logger]
@@ -76,12 +73,21 @@ async def test_generic_api_callback():
# Get the actual request body from the mock
actual_url = mock_post.call_args[1]["url"]
print("##########\n")
print("logs were flushed to URL", actual_url, "with the following headers", mock_post.call_args[1]["headers"])
assert actual_url == test_endpoint, f"Expected URL {test_endpoint}, got {actual_url}"
print(
"logs were flushed to URL",
actual_url,
"with the following headers",
mock_post.call_args[1]["headers"],
)
assert (
actual_url == test_endpoint
), f"Expected URL {test_endpoint}, got {actual_url}"
# Validate headers
assert mock_post.call_args[1]["headers"]["Content-Type"] == "application/json", "Content-Type should be application/json"
assert (
mock_post.call_args[1]["headers"]["Content-Type"] == "application/json"
), "Content-Type should be application/json"
# For the GenericAPILogger, it sends the payload directly as JSON in the data field
json_data = mock_post.call_args[1]["data"]
# Parse the JSON string
@@ -89,27 +95,30 @@ async def test_generic_api_callback():
print("##########\n")
print("json_data", json_data)
actual_request = json.loads(json_data)
# The payload is a list of StandardLoggingPayload objects in the log queue
assert isinstance(actual_request, list), "Request body should be a list"
assert len(actual_request) > 0, "Request body list should not be empty"
# Validate the first payload item
payload_item: StandardLoggingPayload = StandardLoggingPayload(**actual_request[0])
print("##########\n")
print(json.dumps(payload_item, indent=4))
print("##########\n")
# Basic assertions for standard logging payload
assert payload_item["response_cost"] > 0, "Response cost should be greater than 0"
assert payload_item["model"] == "gpt-4o", "Model should be gpt-4o"
assert payload_item["model_parameters"]["user"] == "test_user", "User should be test_user"
assert (
payload_item["model_parameters"]["user"] == "test_user"
), "User should be test_user"
assert payload_item["model"] == "gpt-4o", "Model should be gpt-4o"
assert payload_item["messages"] == [{"role": "user", "content": "Hello, world!"}], "Messages should be the same"
assert payload_item["response"]["choices"][0]["message"]["content"] == "hi", "Response should be hi"
assert payload_item["messages"] == [
{"role": "user", "content": "Hello, world!"}
], "Messages should be the same"
assert (
payload_item["response"]["choices"][0]["message"]["content"] == "hi"
), "Response should be hi"
@pytest.mark.asyncio
@@ -129,9 +138,7 @@ async def test_generic_api_callback_multiple_logs():
# Initialize the GenericAPILogger and set the mock
generic_logger = GenericAPILogger(
endpoint=test_endpoint,
headers=test_headers,
flush_interval=5
endpoint=test_endpoint, headers=test_headers, flush_interval=5
)
generic_logger.async_httpx_client.post = mock_post
litellm.callbacks = [generic_logger]
@@ -154,9 +161,16 @@ async def test_generic_api_callback_multiple_logs():
# Get the actual request body from the mock
actual_url = mock_post.call_args[1]["url"]
print("##########\n")
print("logs were flushed to URL", actual_url, "with the following headers", mock_post.call_args[1]["headers"])
assert actual_url == test_endpoint, f"Expected URL {test_endpoint}, got {actual_url}"
print(
"logs were flushed to URL",
actual_url,
"with the following headers",
mock_post.call_args[1]["headers"],
)
assert (
actual_url == test_endpoint
), f"Expected URL {test_endpoint}, got {actual_url}"
# For the GenericAPILogger, it sends the payload directly as JSON in the data field
json_data = mock_post.call_args[1]["data"]
# Parse the JSON string
@@ -164,12 +178,14 @@ async def test_generic_api_callback_multiple_logs():
print("##########\n")
print("json_data", json_data)
actual_request = json.loads(json_data)
# The payload is a list of StandardLoggingPayload objects in the log queue
assert isinstance(actual_request, list), "Request body should be a list"
assert len(actual_request) > 0, "Request body list should not be empty"
assert len(actual_request) == 10, "Request body list should be 10 items, since we made 10 calls"
assert (
len(actual_request) == 10
), "Request body list should be 10 items, since we made 10 calls"
# Validate all payload items
for payload_item in actual_request:
payload_item: StandardLoggingPayload = StandardLoggingPayload(**payload_item)
@@ -177,10 +193,17 @@ async def test_generic_api_callback_multiple_logs():
print(json.dumps(payload_item, indent=4))
print("##########\n")
assert payload_item["response_cost"] > 0, "Response cost should be greater than 0"
assert (
payload_item["response_cost"] > 0
), "Response cost should be greater than 0"
assert payload_item["model"] == "gpt-4o", "Model should be gpt-4o"
assert payload_item["model_parameters"]["user"] == "test_user", "User should be test_user"
assert (
payload_item["model_parameters"]["user"] == "test_user"
), "User should be test_user"
assert payload_item["model"] == "gpt-4o", "Model should be gpt-4o"
assert payload_item["messages"] == [{"role": "user", "content": "Hello, world!"}], "Messages should be the same"
assert payload_item["response"]["choices"][0]["message"]["content"] == "hi", "Response should be hi"
assert payload_item["messages"] == [
{"role": "user", "content": "Hello, world!"}
], "Messages should be the same"
assert (
payload_item["response"]["choices"][0]["message"]["content"] == "hi"
), "Response should be hi"