mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-27 01:22:18 +00:00
fix(responses): forward safety_identifier through the chat completion bridge
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
e484a7c89c
commit
83223885e6
6 changed files with 109 additions and 6 deletions
|
|
@ -454,6 +454,7 @@ class LiteLLMCompletionResponsesConfig:
|
|||
"stream": stream,
|
||||
"metadata": kwargs.get("metadata"),
|
||||
"service_tier": kwargs.get("service_tier"),
|
||||
"safety_identifier": responses_api_request.get("safety_identifier"),
|
||||
"web_search_options": web_search_options,
|
||||
"response_format": response_format,
|
||||
"reasoning_effort": reasoning.effort,
|
||||
|
|
|
|||
|
|
@ -75,6 +75,7 @@ class ResponsesRequest(BaseModel):
|
|||
stream: bool = False
|
||||
tools: list[ResponsesFunctionTool] | None = None
|
||||
guardrails: list[str] | None = None
|
||||
safety_identifier: str | None = None
|
||||
cache: dict[str, bool] | None = {"no-cache": True}
|
||||
|
||||
|
||||
|
|
@ -316,6 +317,7 @@ class EndpointsClient:
|
|||
*,
|
||||
stream: bool = False,
|
||||
guardrails: list[str] | None = None,
|
||||
safety_identifier: str | None = None,
|
||||
) -> StreamingResponse:
|
||||
return self._send(
|
||||
"/v1/responses",
|
||||
|
|
@ -326,6 +328,7 @@ class EndpointsClient:
|
|||
instructions="You are a helpful assistant",
|
||||
stream=stream,
|
||||
guardrails=guardrails,
|
||||
safety_identifier=safety_identifier,
|
||||
),
|
||||
stream=stream,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8,13 +8,18 @@ litellm-regression-tests/tests/test_inference_endpoints.py.
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import cast
|
||||
import threading
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from types import MappingProxyType
|
||||
from typing import Final, cast
|
||||
|
||||
import pytest
|
||||
from e2e_config import unique_marker
|
||||
from e2e_config import PROVIDER_EDGE_ADVERTISE_HOST, PROVIDER_EDGE_BIND_HOST, unique_marker
|
||||
from e2e_http import (
|
||||
assert_client_error,
|
||||
require_successful_call,
|
||||
unwrap,
|
||||
)
|
||||
from endpoints_client import (
|
||||
EndpointsClient,
|
||||
|
|
@ -26,7 +31,9 @@ from endpoints_client import (
|
|||
ResponsesStreamEventType,
|
||||
)
|
||||
from lifecycle import ResourceManager
|
||||
from models import LiteLLMParamsBody
|
||||
from models import ChatBody, ChatMessage, LiteLLMParamsBody
|
||||
from provider_edge import LiveEdge, start_provider_edge
|
||||
from provider_edge_bedrock import bedrock_signer
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
|
@ -39,6 +46,33 @@ class _OptionalResponsesBody(BaseModel):
|
|||
|
||||
|
||||
BEDROCK_CONVERSE_BACKEND = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
|
||||
BEDROCK_EDGE_REGION: Final = "us-east-1"
|
||||
BEDROCK_EDGE_MOUNT: Final = f"bedrock/{BEDROCK_EDGE_REGION}"
|
||||
|
||||
|
||||
class ConverseRequestBody(BaseModel):
|
||||
additionalModelRequestFields: dict[str, str] | None = None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ConverseRequestCapture:
|
||||
"""The Converse bodies the proxy actually sent upstream, as seen by a live
|
||||
edge sitting between the proxy and Bedrock."""
|
||||
|
||||
_bodies: list[ConverseRequestBody] = field(default_factory=list)
|
||||
_lock: threading.Lock = field(default_factory=threading.Lock)
|
||||
|
||||
def observe(self, url: str, headers: Mapping[str, str], body: bytes | None) -> None:
|
||||
if body is None or "/converse" not in url:
|
||||
return
|
||||
with self._lock:
|
||||
self._bodies.append(ConverseRequestBody.model_validate_json(body))
|
||||
|
||||
@property
|
||||
def bodies(self) -> tuple[ConverseRequestBody, ...]:
|
||||
with self._lock:
|
||||
return tuple(self._bodies)
|
||||
|
||||
|
||||
WEATHER_TOOL = ResponsesFunctionTool(
|
||||
name="get_weather",
|
||||
|
|
@ -295,6 +329,55 @@ class TestResponses:
|
|||
arguments = WeatherArguments.model_validate(raw_arguments)
|
||||
assert arguments.location, f"function call arguments missing location: {function_call.arguments}"
|
||||
|
||||
@pytest.mark.parametrize("endpoint", ["/v1/responses", "/v1/chat/completions"])
|
||||
def test_bedrock_forwards_allowed_safety_identifier_as_additional_model_request_field(
|
||||
self, endpoints_client: EndpointsClient, resources: ResourceManager, endpoint: str
|
||||
) -> None:
|
||||
capture: Final = ConverseRequestCapture()
|
||||
edge: Final = start_provider_edge(
|
||||
LiveEdge(observe_request=capture.observe, sign=bedrock_signer(BEDROCK_EDGE_REGION)),
|
||||
mounts=MappingProxyType({BEDROCK_EDGE_MOUNT: f"https://bedrock-runtime.{BEDROCK_EDGE_REGION}.amazonaws.com"}),
|
||||
bind_host=PROVIDER_EDGE_BIND_HOST,
|
||||
advertise_host=PROVIDER_EDGE_ADVERTISE_HOST,
|
||||
)
|
||||
resources.defer(edge.shutdown)
|
||||
model: Final = f"e2e-responses-{unique_marker()}"
|
||||
model_id: Final = endpoints_client.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(
|
||||
model=BEDROCK_CONVERSE_BACKEND,
|
||||
api_base=edge.edge.api_base(BEDROCK_EDGE_MOUNT),
|
||||
aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID",
|
||||
aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY",
|
||||
aws_region_name=BEDROCK_EDGE_REGION,
|
||||
allowed_openai_params=["safety_identifier"],
|
||||
),
|
||||
)
|
||||
resources.defer(lambda: endpoints_client.delete_model(model_id))
|
||||
key: Final = resources.key()
|
||||
safety_identifier: Final = f"end-user-{unique_marker()}"
|
||||
|
||||
if endpoint == "/v1/responses":
|
||||
responses_result: Final = endpoints_client.responses(
|
||||
key, model, "reply with one word", safety_identifier=safety_identifier
|
||||
)
|
||||
require_successful_call(responses_result)
|
||||
else:
|
||||
unwrap(
|
||||
endpoints_client.proxy.chat(
|
||||
key,
|
||||
ChatBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content="reply with one word")],
|
||||
safety_identifier=safety_identifier,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
assert [body.additionalModelRequestFields for body in capture.bodies] == [
|
||||
{"safety_identifier": safety_identifier}
|
||||
], f"{endpoint} did not forward safety_identifier to Bedrock Converse: {capture.bodies}"
|
||||
|
||||
@pytest.mark.skip(reason="stage red: product gap, /v1/responses 500s (aresponses TypeError) on missing input instead of 400")
|
||||
@pytest.mark.covers("llm.responses.openai.input_validation.nonstream.works")
|
||||
def test_missing_input_returns_error(
|
||||
|
|
|
|||
|
|
@ -298,6 +298,7 @@ class ChatBody(BaseModel):
|
|||
max_completion_tokens: int | None = None
|
||||
temperature: float | None = None
|
||||
user: str | None = None
|
||||
safety_identifier: str | None = None
|
||||
metadata: ChatMetadata | None = None
|
||||
reasoning_effort: str | None = None
|
||||
thinking: ThinkingParam | None = None
|
||||
|
|
@ -976,6 +977,7 @@ class LiteLLMParamsBody(BaseModel):
|
|||
api_base: str | None = None
|
||||
api_version: str | None = None
|
||||
realtime_protocol: str | None = None
|
||||
allowed_openai_params: list[str] | None = None
|
||||
aws_access_key_id: str | None = None
|
||||
aws_secret_access_key: str | None = None
|
||||
aws_region_name: str | None = None
|
||||
|
|
|
|||
|
|
@ -99,6 +99,7 @@ from provider_cache import (
|
|||
SIGNATURE_HEADERS,
|
||||
CacheEdge,
|
||||
MountPolicy,
|
||||
RequestSigner,
|
||||
is_bedrock,
|
||||
scoped_edge_base,
|
||||
split_test_segment,
|
||||
|
|
@ -539,6 +540,7 @@ class ReplayEdge:
|
|||
@dataclass(frozen=True, slots=True)
|
||||
class LiveEdge:
|
||||
observe_request: Callable[[str, Mapping[str, str], bytes | None], None] | None = None
|
||||
sign: RequestSigner | None = None
|
||||
|
||||
|
||||
type EdgeBackend = RecordEdge | ReplayEdge | LiveEdge | CacheEdge
|
||||
|
|
@ -788,14 +790,16 @@ def _handle_live(
|
|||
method: str, url: str, headers: Mapping[str, str], body: bytes | None, timeout: float,
|
||||
cache: CacheEdge | None = None, mount: str = "", test_key: str | None = None,
|
||||
observe_request: Callable[[str, Mapping[str, str], bytes | None], None] | None = None,
|
||||
sign: RequestSigner | None = None,
|
||||
) -> EdgeOutcome:
|
||||
forwarded: Final = {
|
||||
name: value for name, value in headers.items() if name.lower() not in _REQUEST_DROPPED_HEADERS
|
||||
}
|
||||
if observe_request is not None:
|
||||
observe_request(url, forwarded, body)
|
||||
outbound: Final = forwarded if sign is None else sign(method, url, forwarded, body)
|
||||
head: Final = (
|
||||
forward_stream(method, url, headers=forwarded, body=body, timeout=timeout)
|
||||
forward_stream(method, url, headers=outbound, body=body, timeout=timeout)
|
||||
if cache is None else cache.forward(mount, method, url, forwarded, body, timeout, test_key=test_key)
|
||||
)
|
||||
match head:
|
||||
|
|
@ -871,10 +875,10 @@ def handle_edge_request(
|
|||
method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout,
|
||||
backend, mount, test_key,
|
||||
)
|
||||
case LiveEdge(observe_request=observe_request):
|
||||
case LiveEdge(observe_request=observe_request, sign=sign):
|
||||
return _handle_live(
|
||||
method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout,
|
||||
observe_request=observe_request,
|
||||
observe_request=observe_request, sign=sign,
|
||||
)
|
||||
case RecordEdge():
|
||||
return _handle_record(
|
||||
|
|
|
|||
|
|
@ -1248,6 +1248,16 @@ class TestFunctionCallTransformation:
|
|||
assert "tool_choice" not in result
|
||||
assert "tools" not in result
|
||||
|
||||
def test_safety_identifier_forwarded_to_chat_completion_request(self) -> None:
|
||||
result: Final = LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request(
|
||||
model="bedrock/global.openai.gpt-5.6-luna",
|
||||
input="hi",
|
||||
responses_api_request={"safety_identifier": "user-7f3a"},
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
|
||||
assert result["safety_identifier"] == "user-7f3a"
|
||||
|
||||
def test_parallel_tool_calls_dropped_when_no_chat_tools_remain(self) -> None:
|
||||
transform: Final = LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request
|
||||
codex_tool_search: Final = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue