mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
Enforce WebRTC transcription model scope
This commit is contained in:
parent
2961910849
commit
9b0602f019
2 changed files with 294 additions and 16 deletions
|
|
@ -27,6 +27,136 @@ from litellm.types.realtime import (
|
|||
router = APIRouter()
|
||||
|
||||
_REALTIME_TOKEN_VERSION = "realtime_v1"
|
||||
_DEFAULT_REALTIME_MODEL = "gpt-4o-realtime-preview"
|
||||
_DEFAULT_TRANSCRIPTION_MODEL = "gpt-realtime-whisper"
|
||||
_ALLOWED_SESSION_TYPES = ("realtime", "transcription")
|
||||
|
||||
|
||||
def _coerce_realtime_session_type(session_type: Optional[str]) -> str:
|
||||
if session_type in _ALLOWED_SESSION_TYPES:
|
||||
return session_type
|
||||
return "realtime"
|
||||
|
||||
|
||||
def _append_model_candidate(candidates: list[str], model: Any) -> None:
|
||||
if isinstance(model, str) and model and model not in candidates:
|
||||
candidates.append(model)
|
||||
|
||||
|
||||
def _transcription_model_candidates_from_session(session: dict) -> list[str]:
|
||||
candidates: list[str] = []
|
||||
|
||||
audio = session.get("audio")
|
||||
if isinstance(audio, dict):
|
||||
audio_input = audio.get("input")
|
||||
if isinstance(audio_input, dict):
|
||||
nested_transcription = audio_input.get("transcription")
|
||||
if isinstance(nested_transcription, dict):
|
||||
_append_model_candidate(
|
||||
candidates,
|
||||
nested_transcription.get("model"),
|
||||
)
|
||||
|
||||
flat_transcription = session.get("input_audio_transcription")
|
||||
if isinstance(flat_transcription, dict):
|
||||
_append_model_candidate(candidates, flat_transcription.get("model"))
|
||||
|
||||
return candidates
|
||||
|
||||
|
||||
def _set_transcription_model_on_session(
|
||||
session: dict,
|
||||
model: str,
|
||||
create_if_missing: bool = False,
|
||||
) -> None:
|
||||
updated_existing_config = False
|
||||
|
||||
flat_transcription = session.get("input_audio_transcription")
|
||||
if isinstance(flat_transcription, dict):
|
||||
session["input_audio_transcription"] = {
|
||||
**flat_transcription,
|
||||
"model": model,
|
||||
}
|
||||
updated_existing_config = True
|
||||
|
||||
audio = session.get("audio")
|
||||
if isinstance(audio, dict):
|
||||
audio_input = audio.get("input")
|
||||
if isinstance(audio_input, dict):
|
||||
nested_transcription = audio_input.get("transcription")
|
||||
if isinstance(nested_transcription, dict):
|
||||
session["audio"] = {
|
||||
**audio,
|
||||
"input": {
|
||||
**audio_input,
|
||||
"transcription": {
|
||||
**nested_transcription,
|
||||
"model": model,
|
||||
},
|
||||
},
|
||||
}
|
||||
updated_existing_config = True
|
||||
|
||||
if updated_existing_config or not create_if_missing:
|
||||
return
|
||||
|
||||
audio = audio if isinstance(audio, dict) else {}
|
||||
audio_input = audio.get("input")
|
||||
audio_input = audio_input if isinstance(audio_input, dict) else {}
|
||||
session["audio"] = {
|
||||
**audio,
|
||||
"input": {
|
||||
**audio_input,
|
||||
"transcription": {"model": model},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
async def _prepare_client_secret_session(
|
||||
req: RealtimeClientSecretRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
llm_model_list: Optional[list],
|
||||
llm_router: Any,
|
||||
) -> tuple[str, Optional[dict], str]:
|
||||
session_type = _coerce_realtime_session_type(
|
||||
req.session.type if req.session else None
|
||||
)
|
||||
session_data: Optional[dict] = (
|
||||
req.session.model_dump(exclude_none=True) if req.session else None
|
||||
)
|
||||
if session_data is not None:
|
||||
session_data["type"] = session_type
|
||||
|
||||
session_model = req.session.model if req.session else None
|
||||
model: str = session_model or req.model or _DEFAULT_REALTIME_MODEL
|
||||
if session_type != "transcription":
|
||||
return model, session_data, session_type
|
||||
|
||||
transcription_model_candidates = _transcription_model_candidates_from_session(
|
||||
session_data or {}
|
||||
)
|
||||
if not transcription_model_candidates:
|
||||
_append_model_candidate(transcription_model_candidates, session_model)
|
||||
_append_model_candidate(transcription_model_candidates, req.model)
|
||||
if not transcription_model_candidates:
|
||||
transcription_model_candidates.append(_DEFAULT_TRANSCRIPTION_MODEL)
|
||||
|
||||
model = transcription_model_candidates[0]
|
||||
for transcription_model in transcription_model_candidates:
|
||||
await can_key_call_resolved_model(
|
||||
model=transcription_model,
|
||||
valid_token=user_api_key_dict,
|
||||
llm_model_list=llm_model_list,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
if session_data is not None:
|
||||
_set_transcription_model_on_session(
|
||||
session=session_data,
|
||||
model=model,
|
||||
create_if_missing=True,
|
||||
)
|
||||
session_data.pop("model", None)
|
||||
return model, session_data, session_type
|
||||
|
||||
|
||||
def _encode_realtime_token_payload(
|
||||
|
|
@ -99,6 +229,7 @@ async def create_realtime_client_secret(
|
|||
add_litellm_data_to_request,
|
||||
general_settings,
|
||||
llm_router,
|
||||
llm_model_list,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
route_request,
|
||||
|
|
@ -111,17 +242,18 @@ async def create_realtime_client_secret(
|
|||
body = await _read_request_body(request=request)
|
||||
req = RealtimeClientSecretRequest(**body)
|
||||
|
||||
model: str = (
|
||||
(req.session.model if req.session else None)
|
||||
or req.model
|
||||
or "gpt-4o-realtime-preview"
|
||||
model, session_data, session_type = await _prepare_client_secret_session(
|
||||
req=req,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
llm_model_list=llm_model_list,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
data = {"model": model}
|
||||
|
||||
# If session is provided, use it; otherwise create one from model
|
||||
if req.session:
|
||||
data["session"] = req.session.model_dump(exclude_none=True)
|
||||
if session_data is not None:
|
||||
data["session"] = session_data
|
||||
elif req.model:
|
||||
# User provided model at root level, convert to session format
|
||||
data["session"] = {"type": "realtime", "model": model}
|
||||
|
|
@ -166,6 +298,8 @@ async def create_realtime_client_secret(
|
|||
"litellm.proxy.realtime_endpoints.webrtc.create_realtime_client_secret(): Exception - %s",
|
||||
str(e),
|
||||
)
|
||||
if isinstance(e, ProxyException):
|
||||
raise e
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
|
|
@ -204,9 +338,7 @@ async def create_realtime_client_secret(
|
|||
user_id=getattr(user_api_key_dict, "user_id", None),
|
||||
team_id=getattr(user_api_key_dict, "team_id", None),
|
||||
expires_at=expires_at if isinstance(expires_at, int) else None,
|
||||
session_type=(
|
||||
req.session.type if req.session and req.session.type else "realtime"
|
||||
),
|
||||
session_type=session_type,
|
||||
)
|
||||
encrypted_token: str = encrypt_value_helper(token_payload)
|
||||
upstream_json["value"] = encrypted_token
|
||||
|
|
@ -287,17 +419,17 @@ async def proxy_realtime_calls(
|
|||
model = (
|
||||
decoded_payload.get("model_id")
|
||||
or request.query_params.get("model")
|
||||
or "gpt-4o-realtime-preview"
|
||||
or _DEFAULT_REALTIME_MODEL
|
||||
)
|
||||
user_id = decoded_payload.get("user_id") or None
|
||||
team_id = decoded_payload.get("team_id") or None
|
||||
session_type = decoded_payload.get("session_type") or "realtime"
|
||||
if session_type not in ("realtime", "transcription"):
|
||||
session_type = "realtime"
|
||||
session_type = _coerce_realtime_session_type(
|
||||
decoded_payload.get("session_type")
|
||||
)
|
||||
else:
|
||||
# Backward compatibility: older tokens contained only encrypted upstream key.
|
||||
openai_ephemeral_key = decrypted_token_value
|
||||
model = request.query_params.get("model", "gpt-4o-realtime-preview")
|
||||
model = request.query_params.get("model", _DEFAULT_REALTIME_MODEL)
|
||||
user_id = None
|
||||
team_id = None
|
||||
session_type = "realtime"
|
||||
|
|
@ -311,11 +443,17 @@ async def proxy_realtime_calls(
|
|||
|
||||
data: dict = {}
|
||||
try:
|
||||
# Build session config for the multipart form data
|
||||
session_config = {
|
||||
"type": session_type,
|
||||
"model": model,
|
||||
}
|
||||
if session_type == "transcription":
|
||||
_set_transcription_model_on_session(
|
||||
session=session_config,
|
||||
model=model,
|
||||
create_if_missing=True,
|
||||
)
|
||||
else:
|
||||
session_config["model"] = model
|
||||
|
||||
data = {
|
||||
"model": model,
|
||||
|
|
|
|||
|
|
@ -241,6 +241,142 @@ async def test_client_secrets_success_with_mock(
|
|||
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_secrets_transcription_rejects_disallowed_nested_model(
|
||||
proxy_app,
|
||||
):
|
||||
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="test-user",
|
||||
models=["gpt-4o-realtime-preview"],
|
||||
)
|
||||
try:
|
||||
client = TestClient(proxy_app, raise_server_exceptions=False)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.route_request") as mock_route_request,
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging,
|
||||
):
|
||||
mock_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
response = client.post(
|
||||
"/v1/realtime/client_secrets",
|
||||
headers={"Authorization": "Bearer sk-test-master-key"},
|
||||
json={
|
||||
"model": "gpt-4o-realtime-preview",
|
||||
"session": {
|
||||
"type": "transcription",
|
||||
"model": "gpt-4o-realtime-preview",
|
||||
"audio": {
|
||||
"input": {
|
||||
"transcription": {
|
||||
"model": "gpt-realtime-whisper"
|
||||
}
|
||||
}
|
||||
},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 403
|
||||
assert "Tried to access gpt-realtime-whisper" in response.text
|
||||
mock_route_request.assert_not_called()
|
||||
finally:
|
||||
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_secrets_transcription_routes_on_nested_model(
|
||||
proxy_app,
|
||||
mock_add_litellm_data,
|
||||
mock_pre_call_hook,
|
||||
):
|
||||
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="test-user",
|
||||
models=["gpt-4o-realtime-preview", "gpt-realtime-whisper"],
|
||||
)
|
||||
captured = {}
|
||||
future_expires_at = int(time.time()) + 3600
|
||||
|
||||
async def _capturing_route(*args, **kwargs):
|
||||
captured["data"] = kwargs.get("data")
|
||||
|
||||
async def _inner():
|
||||
resp = MagicMock(spec=httpx.Response)
|
||||
resp.status_code = 200
|
||||
resp.text = (
|
||||
f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}'
|
||||
)
|
||||
resp.content = (
|
||||
f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}'
|
||||
).encode()
|
||||
resp.headers = {}
|
||||
resp.json.return_value = {
|
||||
"value": "upstream_ephemeral_key",
|
||||
"expires_at": future_expires_at,
|
||||
}
|
||||
return resp
|
||||
|
||||
return _inner()
|
||||
|
||||
try:
|
||||
client = TestClient(proxy_app)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.route_request",
|
||||
side_effect=_capturing_route,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.add_litellm_data_to_request",
|
||||
side_effect=mock_add_litellm_data,
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_logging,
|
||||
):
|
||||
mock_logging.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook)
|
||||
mock_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
response = client.post(
|
||||
"/v1/realtime/client_secrets",
|
||||
headers={"Authorization": "Bearer sk-test-master-key"},
|
||||
json={
|
||||
"model": "gpt-4o-realtime-preview",
|
||||
"session": {
|
||||
"type": "transcription",
|
||||
"model": "gpt-4o-realtime-preview",
|
||||
"audio": {
|
||||
"input": {
|
||||
"transcription": {
|
||||
"model": "gpt-realtime-whisper"
|
||||
}
|
||||
}
|
||||
},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert captured["data"]["model"] == "gpt-realtime-whisper"
|
||||
session = captured["data"]["session"]
|
||||
assert session["type"] == "transcription"
|
||||
assert "model" not in session
|
||||
assert (
|
||||
session["audio"]["input"]["transcription"]["model"]
|
||||
== "gpt-realtime-whisper"
|
||||
)
|
||||
encrypted_value = response.json()["value"]
|
||||
decoded = _decode_realtime_token_payload(
|
||||
decrypt_value_helper(
|
||||
encrypted_value,
|
||||
key="client_secret.value",
|
||||
exception_type="debug",
|
||||
)
|
||||
or ""
|
||||
)
|
||||
assert decoded is not None
|
||||
assert decoded["model_id"] == "gpt-realtime-whisper"
|
||||
assert decoded["session_type"] == "transcription"
|
||||
finally:
|
||||
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_realtime_calls_requires_auth(proxy_app):
|
||||
"""POST /v1/realtime/calls returns 401 without Authorization.
|
||||
|
||||
|
|
@ -384,6 +520,10 @@ async def test_realtime_calls_replays_transcription_session_type(
|
|||
)
|
||||
|
||||
assert captured["session"]["type"] == "transcription"
|
||||
assert (
|
||||
captured["session"]["audio"]["input"]["transcription"]["model"]
|
||||
== "gpt-realtime-whisper"
|
||||
)
|
||||
|
||||
|
||||
# --- transcription_sessions endpoint ---
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue