Merge pull request #41327 from BerriAI/litellm_s3_log_prompts_only

feat(s3): add s3_log_prompts_only option to log prompts without responses
This commit is contained in:
Yassin Kortam 2026-09-16 14:48:18 -07:00 committed by GitHub
commit 314c0d71a5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 495 additions and 29 deletions

View file

@ -54,6 +54,7 @@ S3_BOUNDED_OBJECT_KEY_HEAD_BYTES: Final = 64
S3_PREFIX_DIGEST_CHARS: Final = 16
# s3 allows 2048 bytes of combined metadata headers, which Content-Disposition counts against
MAX_S3_OBJECT_DOWNLOAD_FILENAME_BYTES: Final = 1024
S3_LOG_PROMPTS_ONLY_ENV_VAR: Final = "S3_LOG_PROMPTS_ONLY"
MAX_FILE_LIST_LIMIT: Final = 10000
DEFAULT_SQS_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_SQS_FLUSH_INTERVAL_SECONDS", 10))
DEFAULT_NUM_WORKERS_LITELLM_PROXY: Final = int(os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", 1))

View file

@ -446,6 +446,12 @@
"ui_name": "S3 Path Prefix",
"description": "Path prefix within the bucket for organizing logs",
"required": false
},
"s3_log_prompts_only": {
"type": "boolean",
"ui_name": "Log Prompts Only",
"description": "Log request messages to S3 but drop the model response from each logged object",
"required": false
}
},
"description": "S3 Bucket (AWS) Logging Integration"

View file

@ -118,6 +118,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
alias_map: Final = {
"langfuse_otel": "langfuse",
"s3_v2": "s3",
}
lookup_name: Final = alias_map.get(normalized_name, normalized_name)

View file

@ -2,19 +2,42 @@
# On success + failure, log events to Supabase
import hashlib
import os
from collections.abc import Mapping
from datetime import datetime
from typing import Final, cast
from pydantic import TypeAdapter, ValidationError
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.constants import (
MAX_S3_OBJECT_DOWNLOAD_FILENAME_BYTES,
MAX_S3_OBJECT_KEY_BYTES,
S3_BOUNDED_OBJECT_KEY_HEAD_BYTES,
S3_LOG_PROMPTS_ONLY_ENV_VAR,
S3_PREFIX_DIGEST_CHARS,
)
from litellm.types.utils import StandardLoggingPayload
_S3_LOG_PROMPTS_ONLY: Final = TypeAdapter(bool)
def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] | None = None) -> bool:
env: Final = os.environ if environ is None else environ
raw: Final = env.get(S3_LOG_PROMPTS_ONLY_ENV_VAR) if configured is None else configured
if raw is None or raw == "":
return False
try:
return _S3_LOG_PROMPTS_ONLY.validate_python(raw.strip() if isinstance(raw, str) else raw)
except ValidationError:
verbose_logger.warning("s3 logging: s3_log_prompts_only=%r is not a boolean, logging prompts only", raw)
return True
def prompts_only_payload(payload: StandardLoggingPayload) -> StandardLoggingPayload:
return {**payload, "response": None}
class S3Logger:
# Class variables or attributes
@ -33,6 +56,7 @@ class S3Logger:
s3_config=None,
s3_server_side_encryption: str | None = None,
s3_sse_kms_key_id: str | None = None,
s3_log_prompts_only: bool | None = None,
**kwargs,
):
import boto3
@ -41,29 +65,30 @@ class S3Logger:
verbose_logger.debug("in init s3 logger - s3_callback_params %s", litellm.s3_callback_params)
s3_use_team_prefix = False
params: Final = {
key: litellm.get_secret(value) if isinstance(value, str) and value.startswith("os.environ/") else value
for key, value in (litellm.s3_callback_params or {}).items()
}
if litellm.s3_callback_params is not None:
# read in .env variables - example os.environ/AWS_BUCKET_NAME
for key, value in litellm.s3_callback_params.items():
if isinstance(value, str) and value.startswith("os.environ/"):
litellm.s3_callback_params[key] = litellm.get_secret(value)
# now set s3 params from litellm.s3_logger_params
s3_bucket_name = litellm.s3_callback_params.get("s3_bucket_name")
s3_region_name = litellm.s3_callback_params.get("s3_region_name")
s3_api_version = litellm.s3_callback_params.get("s3_api_version")
s3_use_ssl = litellm.s3_callback_params.get("s3_use_ssl", True)
s3_verify = litellm.s3_callback_params.get("s3_verify")
s3_endpoint_url = litellm.s3_callback_params.get("s3_endpoint_url")
s3_aws_access_key_id = litellm.s3_callback_params.get("s3_aws_access_key_id")
s3_aws_secret_access_key = litellm.s3_callback_params.get("s3_aws_secret_access_key")
s3_aws_session_token = litellm.s3_callback_params.get("s3_aws_session_token")
s3_config = litellm.s3_callback_params.get("s3_config")
s3_path = litellm.s3_callback_params.get("s3_path")
s3_server_side_encryption = litellm.s3_callback_params.get("s3_server_side_encryption")
s3_sse_kms_key_id = litellm.s3_callback_params.get("s3_sse_kms_key_id")
# done reading litellm.s3_callback_params
s3_use_team_prefix = bool(litellm.s3_callback_params.get("s3_use_team_prefix", False))
s3_bucket_name = params.get("s3_bucket_name")
s3_region_name = params.get("s3_region_name")
s3_api_version = params.get("s3_api_version")
s3_use_ssl = params.get("s3_use_ssl", True)
s3_verify = params.get("s3_verify")
s3_endpoint_url = params.get("s3_endpoint_url")
s3_aws_access_key_id = params.get("s3_aws_access_key_id")
s3_aws_secret_access_key = params.get("s3_aws_secret_access_key")
s3_aws_session_token = params.get("s3_aws_session_token")
s3_config = params.get("s3_config")
s3_path = params.get("s3_path")
s3_server_side_encryption = params.get("s3_server_side_encryption")
s3_sse_kms_key_id = params.get("s3_sse_kms_key_id")
s3_use_team_prefix = bool(params.get("s3_use_team_prefix", False))
self.s3_use_team_prefix = s3_use_team_prefix
self.s3_log_prompts_only: object = (
params.get("s3_log_prompts_only") if s3_log_prompts_only is None else s3_log_prompts_only
)
self.bucket_name = s3_bucket_name
self.s3_path = s3_path
self.s3_server_side_encryption, self.s3_sse_kms_key_id = resolve_sse_params(
@ -144,7 +169,9 @@ class S3Logger:
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
payload_str: Final = safe_dumps(payload)
payload_str: Final = safe_dumps(
prompts_only_payload(payload) if resolve_s3_log_prompts_only(self.s3_log_prompts_only) else payload
)
print_verbose(f"\ns3 Logger - Logging payload = {payload_str}")

View file

@ -21,6 +21,8 @@ from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_S
from litellm.integrations.s3 import (
get_s3_object_download_filename,
get_s3_object_key,
prompts_only_payload,
resolve_s3_log_prompts_only,
resolve_sse_params,
)
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
@ -68,6 +70,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_use_virtual_hosted_style: bool = False,
s3_server_side_encryption: str | None = None,
s3_sse_kms_key_id: str | None = None,
s3_log_prompts_only: bool | None = None,
s3_callback_params_override: dict | None = None,
**kwargs,
):
@ -108,6 +111,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_use_virtual_hosted_style=s3_use_virtual_hosted_style,
s3_server_side_encryption=s3_server_side_encryption,
s3_sse_kms_key_id=s3_sse_kms_key_id,
s3_log_prompts_only=s3_log_prompts_only,
)
verbose_logger.debug("s3 logger using endpoint url %s", s3_endpoint_url)
@ -163,6 +167,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_use_virtual_hosted_style: bool = False,
s3_server_side_encryption: str | None = None,
s3_sse_kms_key_id: str | None = None,
s3_log_prompts_only: bool | None = None,
params_source: dict | None = None,
):
"""
@ -212,6 +217,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
bool(params.get("s3_use_virtual_hosted_style", False)) or s3_use_virtual_hosted_style
)
self.s3_log_prompts_only: object = (
params.get("s3_log_prompts_only") if s3_log_prompts_only is None else s3_log_prompts_only
)
self.s3_server_side_encryption, self.s3_sse_kms_key_id = resolve_sse_params(
params.get("s3_server_side_encryption") or s3_server_side_encryption,
params.get("s3_sse_kms_key_id") or s3_sse_kms_key_id,
@ -489,8 +498,13 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_object_download_filename: Final = get_s3_object_download_filename(start_time, standard_logging_payload["id"])
payload: Final = (
prompts_only_payload(standard_logging_payload)
if resolve_s3_log_prompts_only(self.s3_log_prompts_only)
else standard_logging_payload
)
return s3BatchLoggingElement(
payload=dict(standard_logging_payload),
payload=dict(payload),
s3_object_key=s3_object_key,
s3_object_download_filename=s3_object_download_filename,
)

View file

@ -3710,6 +3710,7 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
"AWS_ACCESS_KEY_ID",
"AWS_SECRET_ACCESS_KEY",
"AWS_REGION_NAME",
"S3_LOG_PROMPTS_ONLY",
],
)

View file

@ -1,16 +1,24 @@
import copy
import json
from datetime import datetime
from unittest.mock import MagicMock, patch
import pytest
import litellm
from litellm.constants import MAX_S3_OBJECT_DOWNLOAD_FILENAME_BYTES, MAX_S3_OBJECT_KEY_BYTES
from litellm.integrations.s3 import S3Logger
from litellm.integrations.s3 import S3Logger, prompts_only_payload, resolve_s3_log_prompts_only
TEST_KMS_KEY_ARN = "arn:aws:kms:us-east-1:111122223333:key/test-key-id"
TEST_MESSAGES = [{"role": "user", "content": "Reply with exactly the word PINEAPPLE."}]
TEST_RESPONSE = {"choices": [{"message": {"role": "assistant", "content": "PINEAPPLE"}}]}
def _standard_logging_payload(response_id: str = "chatcmpl-test-id") -> dict:
return {
"id": response_id,
"messages": copy.deepcopy(TEST_MESSAGES),
"response": copy.deepcopy(TEST_RESPONSE),
"metadata": {"user_api_key_team_alias": None},
}
@ -22,7 +30,9 @@ def _log_event_kwargs(response_id: str = "chatcmpl-test-id") -> dict:
}
def _run_log_event(callback_params: dict, response_id: str = "chatcmpl-test-id") -> MagicMock:
def _run_log_event(
callback_params: dict, response_id: str = "chatcmpl-test-id", log_kwargs: dict[str, object] | None = None
) -> MagicMock:
original = litellm.s3_callback_params
litellm.s3_callback_params = callback_params
try:
@ -31,7 +41,7 @@ def _run_log_event(callback_params: dict, response_id: str = "chatcmpl-test-id")
mock_boto3_client.return_value = mock_s3_client
logger = S3Logger()
logger.log_event(
kwargs=_log_event_kwargs(response_id),
kwargs=_log_event_kwargs(response_id) if log_kwargs is None else log_kwargs,
response_obj={"id": response_id},
start_time=datetime(2026, 7, 30, 12, 0, 0),
end_time=datetime(2026, 7, 30, 12, 0, 1),
@ -182,3 +192,123 @@ def test_put_object_keeps_the_configured_path_intact_when_only_the_id_has_to_shr
key = mock_s3_client.put_object.call_args.kwargs["Key"]
assert key.startswith(long_path + "/2026-07-30/")
assert len(key.encode("utf-8")) == MAX_S3_OBJECT_KEY_BYTES
def _uploaded_body(mock_s3_client: MagicMock) -> dict[str, object]:
return json.loads(mock_s3_client.put_object.call_args.kwargs["Body"])
def test_log_event_prompts_only_drops_response_and_keeps_messages(monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
log_kwargs = _log_event_kwargs()
original_payload = copy.deepcopy(log_kwargs["standard_logging_object"])
mock_s3_client = _run_log_event(
{"s3_bucket_name": "test-bucket", "s3_region_name": "us-east-1", "s3_log_prompts_only": True},
log_kwargs=log_kwargs,
)
body = _uploaded_body(mock_s3_client)
assert body["messages"] == TEST_MESSAGES
assert body["response"] is None
assert body["id"] == "chatcmpl-test-id"
assert log_kwargs["standard_logging_object"] == original_payload
def test_log_event_default_keeps_response(monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
mock_s3_client = _run_log_event({"s3_bucket_name": "test-bucket", "s3_region_name": "us-east-1"})
body = _uploaded_body(mock_s3_client)
assert body["response"] == TEST_RESPONSE
assert body["messages"] == TEST_MESSAGES
def test_log_event_reads_prompts_only_env_var_at_log_time(monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
original = litellm.s3_callback_params
litellm.s3_callback_params = {"s3_bucket_name": "test-bucket", "s3_region_name": "us-east-1"}
try:
with patch("boto3.client") as mock_boto3_client:
mock_s3_client = MagicMock()
mock_boto3_client.return_value = mock_s3_client
logger = S3Logger()
monkeypatch.setenv("S3_LOG_PROMPTS_ONLY", "true")
logger.log_event(
kwargs=_log_event_kwargs(),
response_obj={"id": "chatcmpl-test-id"},
start_time=datetime(2026, 7, 30, 12, 0, 0),
end_time=datetime(2026, 7, 30, 12, 0, 1),
print_verbose=lambda *args, **kwargs: None,
)
finally:
litellm.s3_callback_params = original
body = _uploaded_body(mock_s3_client)
assert body["response"] is None
assert body["messages"] == TEST_MESSAGES
def test_log_event_explicit_false_param_beats_env_var(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("S3_LOG_PROMPTS_ONLY", "true")
mock_s3_client = _run_log_event(
{"s3_bucket_name": "test-bucket", "s3_region_name": "us-east-1", "s3_log_prompts_only": False}
)
assert _uploaded_body(mock_s3_client)["response"] == TEST_RESPONSE
def test_s3_logger_init_does_not_mutate_global_callback_params(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("MY_S3_BUCKET", "resolved-bucket")
callback_params = {"s3_bucket_name": "os.environ/MY_S3_BUCKET", "s3_region_name": "us-east-1"}
snapshot = copy.deepcopy(callback_params)
original = litellm.s3_callback_params
litellm.s3_callback_params = callback_params
try:
with patch("boto3.client"):
logger = S3Logger()
finally:
litellm.s3_callback_params = original
assert logger.bucket_name == "resolved-bucket"
assert callback_params == snapshot
@pytest.mark.parametrize(
"configured,env_value,expected",
[
(True, None, True),
(False, "true", False),
("true", None, True),
("False", "true", False),
("1", None, True),
("0", None, False),
(" yes ", None, True),
(None, None, False),
(None, "true", True),
(None, "false", False),
(None, "", False),
("", "true", False),
],
)
def test_resolve_s3_log_prompts_only(configured: object, env_value: str | None, expected: bool):
environ = {} if env_value is None else {"S3_LOG_PROMPTS_ONLY": env_value}
assert resolve_s3_log_prompts_only(configured, environ) is expected
def test_resolve_s3_log_prompts_only_unparseable_value_fails_toward_prompts_only():
assert resolve_s3_log_prompts_only("enabled", {}) is True
def test_prompts_only_payload_returns_copy_with_response_cleared():
payload = _standard_logging_payload()
snapshot = copy.deepcopy(payload)
stripped = prompts_only_payload(payload)
assert stripped["response"] is None
assert stripped["messages"] == TEST_MESSAGES
assert stripped is not payload
assert payload == snapshot

View file

@ -1,8 +1,11 @@
import asyncio
import copy
import json
import re
import sys
import textwrap
import uuid
from collections.abc import Awaitable, Callable
from contextlib import asynccontextmanager
from datetime import datetime
from pathlib import Path
@ -10,6 +13,7 @@ from unittest.mock import AsyncMock, MagicMock, call, patch
import httpx
import pytest
import respx
from litellm.integrations.s3_v2 import S3Logger
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
@ -2310,3 +2314,137 @@ def _s3_logger_for_region(region_name: str) -> S3Logger:
)
def test_build_object_url_uses_partition_dns_suffix(region_name: str, expected_url: str) -> None:
assert _s3_logger_for_region(region_name)._build_object_url("2025-01-01/key.json") == expected_url
def _prompts_only_logger(s3_log_prompts_only: bool | None = None) -> S3Logger:
return S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_log_prompts_only=s3_log_prompts_only,
)
def _chat_payload() -> StandardLoggingPayload:
return StandardLoggingPayload(
id="chatcmpl-prompts-only",
messages=[{"role": "user", "content": "Reply with exactly the word PINEAPPLE."}],
response={"choices": [{"message": {"role": "assistant", "content": "PINEAPPLE"}}]},
metadata={"user_api_key_team_alias": None},
)
async def _queued_body_via_async_upload(
logger: S3Logger, log_event: Callable[..., Awaitable[None]]
) -> dict[str, object]:
payload = _chat_payload()
original = copy.deepcopy(payload)
await log_event(
kwargs={"standard_logging_object": payload},
response_obj=None,
start_time=datetime(2026, 7, 30, 12, 0, 0),
end_time=datetime(2026, 7, 30, 12, 0, 1),
)
assert payload == original, "the caller's standard_logging_object must not be mutated"
(element,) = logger.log_queue
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put.return_value = response
await logger.async_upload_data_to_s3(element)
return json.loads(logger.async_httpx_client.put.call_args.kwargs["data"])
@pytest.mark.asyncio
@pytest.mark.parametrize("event_name", ["async_log_success_event", "async_log_failure_event"])
async def test_prompts_only_drops_response_but_keeps_messages_in_uploaded_object(
monkeypatch: pytest.MonkeyPatch, event_name: str
):
import litellm
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_log_prompts_only": True})
logger = _prompts_only_logger()
log_event: Callable[..., Awaitable[None]] = (
logger.async_log_success_event if event_name == "async_log_success_event" else logger.async_log_failure_event
)
body = await _queued_body_via_async_upload(logger, log_event)
assert body["messages"] == _chat_payload()["messages"]
assert body["response"] is None
assert body["id"] == "chatcmpl-prompts-only"
@pytest.mark.asyncio
async def test_prompts_only_default_off_keeps_response_in_uploaded_object(monkeypatch: pytest.MonkeyPatch):
import litellm
monkeypatch.setattr(litellm, "s3_callback_params", {})
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
logger = _prompts_only_logger()
body = await _queued_body_via_async_upload(logger, logger.async_log_success_event)
assert body["response"] == _chat_payload()["response"]
assert body["messages"] == _chat_payload()["messages"]
@pytest.mark.asyncio
async def test_prompts_only_explicit_false_in_params_beats_env_var(monkeypatch: pytest.MonkeyPatch):
import litellm
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_log_prompts_only": False})
monkeypatch.setenv("S3_LOG_PROMPTS_ONLY", "true")
logger = _prompts_only_logger()
body = await _queued_body_via_async_upload(logger, logger.async_log_success_event)
assert body["response"] == _chat_payload()["response"]
@pytest.mark.asyncio
async def test_prompts_only_env_var_applies_when_param_unset(monkeypatch: pytest.MonkeyPatch):
import litellm
monkeypatch.setattr(litellm, "s3_callback_params", {})
logger = _prompts_only_logger()
monkeypatch.setenv("S3_LOG_PROMPTS_ONLY", "true")
body = await _queued_body_via_async_upload(logger, logger.async_log_success_event)
assert body["response"] is None
assert body["messages"] == _chat_payload()["messages"]
@respx.mock
def test_prompts_only_constructor_kwarg_applies_to_sync_upload(monkeypatch: pytest.MonkeyPatch):
import litellm
monkeypatch.setattr(litellm, "s3_callback_params", {})
monkeypatch.delenv("S3_LOG_PROMPTS_ONLY", raising=False)
logger = _prompts_only_logger(s3_log_prompts_only=True)
payload = _chat_payload()
element = logger.create_s3_batch_logging_element(
start_time=datetime(2026, 7, 30, 12, 0, 0),
standard_logging_payload=payload,
)
assert element is not None
assert payload["response"] == _chat_payload()["response"]
put_route = respx.put(url__regex=r"https://test-bucket\.s3\..*").mock(return_value=httpx.Response(200))
logger.upload_data_to_s3(element)
body = json.loads(put_route.calls.last.request.content)
assert body["response"] is None
assert body["messages"] == _chat_payload()["messages"]
@pytest.mark.parametrize("callback_name", ["s3", "s3_v2"])
def test_prompts_only_toggle_is_exposed_to_admin_ui_for_both_s3_callbacks(callback_name: str):
from litellm.integrations.custom_logger import CustomLogger
assert "S3_LOG_PROMPTS_ONLY" in CustomLogger.get_callback_env_vars(callback_name)

View file

@ -302,6 +302,113 @@ describe("Settings", () => {
});
});
const mockS3Callback = (variables: Record<string, string | null>, callbackName = "s3") => {
mockGetCallbacksCall.mockResolvedValue({
callbacks: [{ name: callbackName, variables }],
available_callbacks: {
s3: {
litellm_callback_name: "s3",
litellm_callback_params: [
"AWS_ACCESS_KEY_ID",
"AWS_SECRET_ACCESS_KEY",
"AWS_REGION_NAME",
"S3_LOG_PROMPTS_ONLY",
],
ui_callback_name: "s3 Bucket (AWS)",
},
},
alerts: [],
});
mockGetCallbackConfigsCall.mockResolvedValue([
{
id: "s3",
displayName: "S3",
dynamic_params: {
s3_bucket_name: { type: "text", ui_name: "S3 Bucket Name", required: false },
s3_log_prompts_only: { type: "boolean", ui_name: "Log Prompts Only", required: false },
},
},
]);
};
const openS3EditModal = async (callbackName = "s3") => {
const user = userEvent.setup();
render(<Settings {...defaultProps} />);
await user.click(await screen.findByTestId(`callback-actions-${callbackName}-success`));
await user.click(await screen.findByTestId("callback-action-edit"));
return user;
};
it("should render a saved boolean dynamic param as a checked switch and post false when toggled off", async () => {
mockS3Callback({ S3_LOG_PROMPTS_ONLY: "true" });
const user = await openS3EditModal();
const promptsOnlySwitch = await screen.findByRole("switch", { name: "Log Prompts Only" });
expect(promptsOnlySwitch).toBeChecked();
await user.click(promptsOnlySwitch);
expect(promptsOnlySwitch).not.toBeChecked();
await user.click(within(screen.getByRole("dialog")).getByRole("button", { name: "Save Changes" }));
await waitFor(() => {
expect(vi.mocked(setCallbacksCall)).toHaveBeenCalledWith(
"token",
expect.objectContaining({
environment_variables: expect.objectContaining({ callback: "s3", s3_log_prompts_only: "false" }),
}),
);
});
});
it("should render an unset boolean dynamic param as an unchecked switch and post true when toggled on", async () => {
mockS3Callback({ S3_LOG_PROMPTS_ONLY: null });
const user = await openS3EditModal();
const promptsOnlySwitch = await screen.findByRole("switch", { name: "Log Prompts Only" });
expect(promptsOnlySwitch).not.toBeChecked();
await user.click(promptsOnlySwitch);
await user.click(within(screen.getByRole("dialog")).getByRole("button", { name: "Save Changes" }));
await waitFor(() => {
expect(vi.mocked(setCallbacksCall)).toHaveBeenCalledWith(
"token",
expect.objectContaining({
environment_variables: expect.objectContaining({ callback: "s3", s3_log_prompts_only: "true" }),
}),
);
});
});
it.each(["True", "1"])("should render a boolean dynamic param stored as %s as a checked switch", async (stored) => {
mockS3Callback({ S3_LOG_PROMPTS_ONLY: stored });
await openS3EditModal();
expect(await screen.findByRole("switch", { name: "Log Prompts Only" })).toBeChecked();
});
it("should resolve the s3_v2 callback to the s3 dynamic params and post under the s3_v2 name", async () => {
mockS3Callback({ S3_LOG_PROMPTS_ONLY: null }, "s3_v2");
const user = await openS3EditModal("s3_v2");
const promptsOnlySwitch = await screen.findByRole("switch", { name: "Log Prompts Only" });
expect(promptsOnlySwitch).not.toBeChecked();
expect(within(screen.getByRole("dialog")).getByRole("combobox", { name: "Callback" })).toHaveValue("S3");
await user.click(promptsOnlySwitch);
await user.click(within(screen.getByRole("dialog")).getByRole("button", { name: "Save Changes" }));
await waitFor(() => {
expect(vi.mocked(setCallbacksCall)).toHaveBeenCalledWith(
"token",
expect.objectContaining({
environment_variables: expect.objectContaining({ callback: "s3_v2", s3_log_prompts_only: "true" }),
litellm_settings: { success_callback: ["s3_v2"] },
}),
);
});
});
it("should send the typed webhook url for an alert type when the alerting tab is saved", async () => {
const user = userEvent.setup();
render(<Settings {...defaultProps} />);

View file

@ -67,19 +67,20 @@ const DynamicParamsFields: React.FC<DynamicParamsFieldsProps> = ({ params, callb
return null;
}
const callbackConfig = findCallbackConfig(callbackConfigs, selectedCallback);
return (
<div className="space-y-4 mt-6 p-4 bg-muted rounded-lg border">
{params.map((param) => {
const callbackConfig = callbackConfigs.find((config) => config.id === selectedCallback);
const paramConfig = callbackConfig?.dynamic_params?.[param] || {};
const paramType = paramConfig.type || "text";
const fieldLabel = paramConfig.ui_name || param.replace(/_/g, " ").replace(/\b\w/g, (l) => l.toUpperCase());
const isRequired = paramConfig.required || false;
const selectOptions: string[] = Array.isArray(paramConfig.options) ? paramConfig.options : [];
const isSelect = paramType === "select" && selectOptions.length > 0;
const isBoolean = paramType === "boolean";
const fieldId = `${fieldIdPrefix}-${param}`;
const validationRules = isRequired ? { required: `Please enter the ${fieldLabel.toLowerCase()}` } : undefined;
const registration = isSelect ? undefined : register(param, validationRules);
const registration = isSelect || isBoolean ? undefined : register(param, validationRules);
return (
<Field key={param} className="mb-4">
@ -111,7 +112,22 @@ const DynamicParamsFields: React.FC<DynamicParamsFieldsProps> = ({ params, callb
)}
/>
)}
{isBoolean && (
<Controller
control={control}
name={param}
render={({ field }) => (
<Switch
id={fieldId}
checked={/^(true|1)$/i.test(String(field.value ?? ""))}
onCheckedChange={(checked: boolean) => field.onChange(checked ? "true" : "false")}
onBlur={field.onBlur}
/>
)}
/>
)}
{!isSelect &&
!isBoolean &&
(paramType === "password" ? (
<Input
id={fieldId}
@ -162,7 +178,7 @@ export const CallbackSelector: React.FC<CallbackSelectorProps> = ({
}) => {
const { control } = useFormContext<CallbackFormValues>();
const inputId = React.useId();
const selectedConfig = callbackConfigs.find((config) => config.id === selectedCallback) ?? null;
const selectedConfig = findCallbackConfig(callbackConfigs, selectedCallback) ?? null;
return (
<Controller
@ -221,6 +237,31 @@ export const CallbackSelector: React.FC<CallbackSelectorProps> = ({
);
};
const CALLBACK_CONFIG_ALIASES: Record<string, string> = { s3_v2: "s3" };
interface DynamicParamConfig {
type?: string;
ui_name?: string;
required?: boolean;
options?: string[];
}
interface CallbackConfigWithParams {
id: string;
dynamic_params?: Record<string, DynamicParamConfig>;
}
const findCallbackConfig = <T extends { id: string }>(
callbackConfigs: readonly T[],
callbackName: string | null,
): T | undefined => {
if (!callbackName) {
return undefined;
}
const configId = CALLBACK_CONFIG_ALIASES[callbackName] ?? callbackName;
return callbackConfigs.find((config) => config.id === configId);
};
// Shared helper function to get dynamic params for a callback
const getDynamicParamsForCallback = (
callbackName: string | null,
@ -231,7 +272,7 @@ const getDynamicParamsForCallback = (
return fallbackVariables ? Object.keys(fallbackVariables) : [];
}
const callbackConfig = callbackConfigs.find((config) => config.id === callbackName);
const callbackConfig = findCallbackConfig(callbackConfigs, callbackName);
if (callbackConfig?.dynamic_params) {
return Object.keys(callbackConfig.dynamic_params);
}