mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
[Feat]Add support for safety_identifier parameter in chat.completions.create (#14174)
* Add support for safety_identifier parameter in chat.completions.create * make sure param is getting actually passed to the raw api
This commit is contained in:
parent
61b2209827
commit
4adfd18bc6
7 changed files with 76 additions and 6 deletions
|
|
@ -106,6 +106,7 @@ def completion(
|
|||
parallel_tool_calls: Optional[bool] = None,
|
||||
logprobs: Optional[bool] = None,
|
||||
top_logprobs: Optional[int] = None,
|
||||
safety_identifier: Optional[str] = None,
|
||||
deployment_id=None,
|
||||
# soon to be deprecated params by OpenAI
|
||||
functions: Optional[List] = None,
|
||||
|
|
@ -196,6 +197,8 @@ def completion(
|
|||
|
||||
- `top_logprobs`: *int (optional)* - An integer between 0 and 5 specifying the number of most likely tokens to return at each token position, each with an associated log probability. `logprobs` must be set to true if this parameter is used.
|
||||
|
||||
- `safety_identifier`: *string (optional)* - A unique identifier for tracking and managing safety-related requests. This parameter helps with safety monitoring and compliance tracking.
|
||||
|
||||
- `headers`: *dict (optional)* - A dictionary of headers to be sent with the request.
|
||||
|
||||
- `extra_headers`: *dict (optional)* - Alternative to `headers`, used to send extra headers in LLM API request.
|
||||
|
|
|
|||
|
|
@ -14,7 +14,9 @@ DEFAULT_S3_BATCH_SIZE = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512))
|
|||
DEFAULT_SQS_FLUSH_INTERVAL_SECONDS = int(
|
||||
os.getenv("DEFAULT_SQS_FLUSH_INTERVAL_SECONDS", 10)
|
||||
)
|
||||
DEFAULT_NUM_WORKERS_LITELLM_PROXY = int(os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", 4))
|
||||
DEFAULT_NUM_WORKERS_LITELLM_PROXY = int(
|
||||
os.getenv("DEFAULT_NUM_WORKERS_LITELLM_PROXY", 4)
|
||||
)
|
||||
DEFAULT_SQS_BATCH_SIZE = int(os.getenv("DEFAULT_SQS_BATCH_SIZE", 512))
|
||||
SQS_SEND_MESSAGE_ACTION = "SendMessage"
|
||||
SQS_API_VERSION = "2012-11-05"
|
||||
|
|
@ -395,6 +397,7 @@ DEFAULT_CHAT_COMPLETION_PARAM_VALUES = {
|
|||
"reasoning_effort": None,
|
||||
"thinking": None,
|
||||
"web_search_options": None,
|
||||
"safety_identifier": None,
|
||||
}
|
||||
|
||||
openai_compatible_endpoints: List = [
|
||||
|
|
|
|||
|
|
@ -158,6 +158,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
"parallel_tool_calls",
|
||||
"audio",
|
||||
"web_search_options",
|
||||
"safety_identifier",
|
||||
] # works across all models
|
||||
|
||||
model_specific_params = []
|
||||
|
|
|
|||
|
|
@ -357,6 +357,7 @@ async def acompletion(
|
|||
top_logprobs: Optional[int] = None,
|
||||
deployment_id=None,
|
||||
reasoning_effort: Optional[Literal["minimal", "low", "medium", "high"]] = None,
|
||||
safety_identifier: Optional[str] = None,
|
||||
# set api_base, api_version, api_key
|
||||
base_url: Optional[str] = None,
|
||||
api_version: Optional[str] = None,
|
||||
|
|
@ -493,6 +494,7 @@ async def acompletion(
|
|||
"api_key": api_key,
|
||||
"model_list": model_list,
|
||||
"reasoning_effort": reasoning_effort,
|
||||
"safety_identifier": safety_identifier,
|
||||
"extra_headers": extra_headers,
|
||||
"acompletion": True, # assuming this is a required parameter
|
||||
"thinking": thinking,
|
||||
|
|
@ -906,6 +908,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
web_search_options: Optional[OpenAIWebSearchOptions] = None,
|
||||
deployment_id=None,
|
||||
extra_headers: Optional[dict] = None,
|
||||
safety_identifier: Optional[str] = None,
|
||||
# soon to be deprecated params by OpenAI
|
||||
functions: Optional[List] = None,
|
||||
function_call: Optional[str] = None,
|
||||
|
|
@ -1243,6 +1246,7 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
"reasoning_effort": reasoning_effort,
|
||||
"thinking": thinking,
|
||||
"web_search_options": web_search_options,
|
||||
"safety_identifier": safety_identifier,
|
||||
"allowed_openai_params": kwargs.get("allowed_openai_params"),
|
||||
}
|
||||
optional_params = get_optional_params(
|
||||
|
|
|
|||
|
|
@ -788,6 +788,7 @@ class ChatCompletionRequest(TypedDict, total=False):
|
|||
response_format: dict
|
||||
seed: int
|
||||
service_tier: str
|
||||
safety_identifier: str
|
||||
stop: Union[str, List[str]]
|
||||
stream_options: dict
|
||||
temperature: float
|
||||
|
|
|
|||
|
|
@ -837,15 +837,13 @@ async def _client_async_logging_helper(
|
|||
# Async Logging Worker
|
||||
################################################
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
async_coroutine = logging_obj.async_success_handler(
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time
|
||||
async_coroutine=logging_obj.async_success_handler(
|
||||
result=result, start_time=start_time, end_time=end_time
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
################################################
|
||||
# Sync Logging Worker
|
||||
################################################
|
||||
|
|
@ -3304,6 +3302,7 @@ def get_optional_params( # noqa: PLR0915
|
|||
messages: Optional[List[AllMessageValues]] = None,
|
||||
thinking: Optional[AnthropicThinkingParam] = None,
|
||||
web_search_options: Optional[OpenAIWebSearchOptions] = None,
|
||||
safety_identifier: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
passed_params = locals().copy()
|
||||
|
|
|
|||
|
|
@ -664,3 +664,62 @@ async def test_openai_gpt5_reasoning():
|
|||
)
|
||||
print("response: ", response)
|
||||
assert response.choices[0].message.content is not None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_safety_identifier_parameter():
|
||||
"""Test that safety_identifier parameter is correctly passed to the OpenAI API."""
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
litellm.set_verbose = True
|
||||
client = AsyncOpenAI(api_key="fake-api-key")
|
||||
|
||||
with patch.object(
|
||||
client.chat.completions.with_raw_response, "create"
|
||||
) as mock_client:
|
||||
try:
|
||||
await litellm.acompletion(
|
||||
model="openai/gpt-4o",
|
||||
messages=[{"role": "user", "content": "Hello, how are you?"}],
|
||||
safety_identifier="user_code_123456",
|
||||
client=client,
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
|
||||
mock_client.assert_called_once()
|
||||
request_body = mock_client.call_args.kwargs
|
||||
|
||||
# Verify the request contains the safety_identifier parameter
|
||||
assert "safety_identifier" in request_body
|
||||
# Verify safety_identifier is correctly sent to the API
|
||||
assert request_body["safety_identifier"] == "user_code_123456"
|
||||
|
||||
|
||||
def test_openai_safety_identifier_parameter_sync():
|
||||
"""Test that safety_identifier parameter is correctly passed to the OpenAI API."""
|
||||
from openai import OpenAI
|
||||
|
||||
litellm.set_verbose = True
|
||||
client = OpenAI(api_key="fake-api-key")
|
||||
|
||||
with patch.object(
|
||||
client.chat.completions.with_raw_response, "create"
|
||||
) as mock_client:
|
||||
try:
|
||||
litellm.completion(
|
||||
model="openai/gpt-4o",
|
||||
messages=[{"role": "user", "content": "Hello, how are you?"}],
|
||||
safety_identifier="user_code_123456",
|
||||
client=client,
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"Error: {e}")
|
||||
|
||||
mock_client.assert_called_once()
|
||||
request_body = mock_client.call_args.kwargs
|
||||
|
||||
# Verify the request contains the safety_identifier parameter
|
||||
assert "safety_identifier" in request_body
|
||||
# Verify safety_identifier is correctly sent to the API
|
||||
assert request_body["safety_identifier"] == "user_code_123456"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue