[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:
Sameer Kankute 2025-09-02 22:07:08 +05:30 • committed by GitHub
parent 61b2209827
commit 4adfd18bc6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 76 additions and 6 deletions

View file

@ -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.

View file

@ -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 = [

View file

@ -158,6 +158,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
"parallel_tool_calls",
"audio",
"web_search_options",
"safety_identifier",
] # works across all models
model_specific_params = []

View file

@ -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(

View file

@ -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

View file

@ -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()

View file

@ -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"