mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(live): keep Live authentication on the model it dispatches
Veria flagged three ways a Live request could be judged against a model other
than the one it runs:
- the synthetic authentication request copied `parsed_body` from its source
scope, so a reauthentication re-read the `{}` body cached by the first
authentication and the per-model budget gate saw no model at all;
- a key with `budget_fallbacks` could be rerouted by the budget check while
`_create()` still dispatched the requested session model, so those requests
now carry `litellm_pinned_realtime_model`, the marker upstream uses to make
that fallback fail closed;
- `image_edit()` spread caller-controlled passthrough parameters over the
authenticated model, letting a restricted key name another image model in
`extra_body`; the authenticated model now wins in both branches.
This commit is contained in:
parent
5110d61901
commit
a3f10ec09c
4 changed files with 87 additions and 8 deletions
|
|
@ -136,9 +136,9 @@ class ChatGPTImageEditConfig(OpenAIImageEditConfig):
|
|||
if not 1 <= len(validated) <= 5:
|
||||
raise ValueError("images must contain between 1 and 5 reference images")
|
||||
return { # mutable-ok: JSON request serialization
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
**image_edit_optional_request_params,
|
||||
"model": model, # the authenticated alias wins over passthrough fields
|
||||
"images": tuple(item.model_dump() for item in validated),
|
||||
}, ()
|
||||
|
||||
|
|
@ -147,8 +147,8 @@ class ChatGPTImageEditConfig(OpenAIImageEditConfig):
|
|||
if not 1 <= len(encoded) <= 5:
|
||||
raise ValueError("images must contain between 1 and 5 reference images")
|
||||
return { # mutable-ok: JSON request serialization
|
||||
"model": model,
|
||||
"prompt": prompt,
|
||||
**image_edit_optional_request_params,
|
||||
"model": model, # the authenticated alias wins over passthrough fields
|
||||
"images": encoded,
|
||||
}, ()
|
||||
|
|
|
|||
|
|
@ -214,13 +214,26 @@ async def _body(request: Request) -> Mapping[str, JsonValue]:
|
|||
raise HTTPException(400, "Expected a JSON object") from exc
|
||||
|
||||
|
||||
def _request(source: Request | WebSocket, body: Mapping[str, JsonValue]) -> Request:
|
||||
def _request(source: Request | WebSocket, body: Mapping[str, JsonValue], pinned_model: str | None = None) -> Request:
|
||||
"""Build the synthetic POST request that authenticates one Live operation.
|
||||
|
||||
A copied scope can carry the body a previous authentication parsed, so the
|
||||
cached ``parsed_body`` is dropped and the request parses ``body`` again.
|
||||
``litellm_pinned_realtime_model`` marks the requests whose model this endpoint
|
||||
dispatches itself, so a key-level budget fallback cannot authorize a different
|
||||
model than the one that is about to run.
|
||||
"""
|
||||
|
||||
async def receive() -> Message:
|
||||
return _mutable(
|
||||
MappingProxyType({"type": "http.request", "body": _encode_json(body).encode(), "more_body": False})
|
||||
)
|
||||
|
||||
return Request(_mutable(MappingProxyType({**source.scope, "type": "http", "method": "POST"})), receive=receive)
|
||||
scope: Final = _mutable(MappingProxyType({**source.scope, "type": "http", "method": "POST"}))
|
||||
scope.pop("parsed_body", None)
|
||||
if pinned_model is not None:
|
||||
scope["litellm_pinned_realtime_model"] = pinned_model
|
||||
return Request(scope, receive=receive)
|
||||
|
||||
|
||||
def _session_model(body: Mapping[str, JsonValue], fallback: str | None = None) -> str:
|
||||
|
|
@ -449,7 +462,7 @@ async def _budget_scope(auth: UserAPIKeyAuth) -> AsyncGenerator[_BudgetOwnership
|
|||
|
||||
async def _reauth(ownership: _BudgetOwnership, request: Request, body: Mapping[str, JsonValue], model: str) -> None:
|
||||
await release_or_invalidate_budget_reservation(budget_reservation=ownership.auth.budget_reservation)
|
||||
ownership.replace_auth(await _auth(_request(request, MappingProxyType({**body, "model": model}))))
|
||||
ownership.replace_auth(await _auth(_request(request, MappingProxyType({**body, "model": model}), model)))
|
||||
|
||||
|
||||
def _policy_object(value: object) -> Mapping[str, JsonValue]:
|
||||
|
|
@ -1288,9 +1301,9 @@ def _response(response: httpx.Response, handle: LiveHandle | None = None) -> Res
|
|||
|
||||
async def _create(request: Request, token: str | None = None) -> Response:
|
||||
body: Final = await _body(request)
|
||||
auth: Final = await _auth(
|
||||
_request(request, _EMPTY if token else MappingProxyType({**body, "model": _session_model(body)}))
|
||||
)
|
||||
requested: Final = None if token else _session_model(body)
|
||||
auth_body: Final = _EMPTY if token or requested is None else MappingProxyType({**body, "model": requested})
|
||||
auth: Final = await _auth(_request(request, auth_body, requested))
|
||||
async with _budget_scope(auth) as ownership:
|
||||
source: Final = decode_session(token, _owner(auth)) if token else None
|
||||
model: Final = _session_model(body, source.alias if source else None)
|
||||
|
|
|
|||
|
|
@ -165,9 +165,11 @@ def route_client(monkeypatch):
|
|||
selected = AsyncMock(return_value=deployment)
|
||||
supervised = AsyncMock()
|
||||
authenticated_bodies = []
|
||||
authenticated_scopes = []
|
||||
|
||||
async def authenticate(request):
|
||||
authenticated_bodies.append(await request.json())
|
||||
authenticated_scopes.append(request.scope)
|
||||
return auth
|
||||
|
||||
@asynccontextmanager
|
||||
|
|
@ -205,6 +207,7 @@ def route_client(monkeypatch):
|
|||
auth=auth,
|
||||
factory=factory,
|
||||
bodies=authenticated_bodies,
|
||||
scopes=authenticated_scopes,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -237,6 +240,47 @@ def test_create_preserves_configuration_and_returns_owned_json_session(route_cli
|
|||
assert route_client.factory.call_args.args[0].api_base is None
|
||||
|
||||
|
||||
def test_synthetic_live_request_drops_the_cached_body_and_pins_the_model_on_request():
|
||||
source = Request({"type": "http", "method": "POST", "headers": [], "path": "/v1/live/sessions"})
|
||||
source.scope["parsed_body"] = (("model",), {"model": "stale-alias"})
|
||||
|
||||
pinned = live._request(source, MappingProxyType({"model": "voice"}), "voice")
|
||||
assert "parsed_body" not in pinned.scope
|
||||
assert pinned.scope["litellm_pinned_realtime_model"] == "voice"
|
||||
|
||||
plain = live._request(source, MappingProxyType({"model": "voice"}))
|
||||
assert "litellm_pinned_realtime_model" not in plain.scope
|
||||
assert "parsed_body" not in plain.scope
|
||||
assert source.scope["parsed_body"] == (("model",), {"model": "stale-alias"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_synthetic_live_request_sends_the_replaced_body():
|
||||
source = Request({"type": "http", "method": "POST", "headers": [], "path": "/v1/live/sessions"})
|
||||
source.scope["parsed_body"] = (("model",), {"model": "stale-alias"})
|
||||
|
||||
request = live._request(source, MappingProxyType({"model": "voice"}))
|
||||
|
||||
assert await request.json() == {"model": "voice"}
|
||||
|
||||
|
||||
def test_live_create_and_fork_pin_the_model_they_dispatch(route_client):
|
||||
created = route_client.client.post(
|
||||
"/v1/live/sessions",
|
||||
json={"session": {"model": "voice"}, "transport": {"type": "webrtc", "sdp": "offer"}},
|
||||
)
|
||||
assert created.status_code == 201
|
||||
assert [scope.get("litellm_pinned_realtime_model") for scope in route_client.scopes] == ["voice"]
|
||||
|
||||
token = live.encode_session(handle(model_id="deployment-a"))
|
||||
forked = route_client.client.post(
|
||||
f"/v1/live/sessions/{token}/fork",
|
||||
json={"session": {}, "transport": {"type": "webrtc", "sdp": "offer"}},
|
||||
)
|
||||
assert forked.status_code == 201
|
||||
assert [scope.get("litellm_pinned_realtime_model") for scope in route_client.scopes[1:]] == [None, "voice"]
|
||||
|
||||
|
||||
def test_fork_preserves_empty_overrides_and_pins_source_deployment(route_client):
|
||||
source = handle(model_id="deployment-a")
|
||||
token = live.encode_session(source)
|
||||
|
|
|
|||
|
|
@ -201,6 +201,28 @@ async def test_async_codex_edit_without_multipart_image(chatgpt_tokens):
|
|||
assert requests[0].headers["x-gateway-route"] == "images"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"references",
|
||||
[None, [{"image_url": "data:image/png;base64,aGVsbG8="}]],
|
||||
ids=["uploaded-image", "reference-images"],
|
||||
)
|
||||
def test_edit_keeps_the_authenticated_model_over_passthrough_fields(tmp_path, references):
|
||||
image = None
|
||||
if references is None:
|
||||
image = tmp_path / "reference.png"
|
||||
image.write_bytes(b"reference image bytes")
|
||||
data, _ = ChatGPTImageEditConfig().transform_image_edit_request(
|
||||
"gpt-image-2",
|
||||
"edit",
|
||||
image,
|
||||
{"model": "gpt-image-2.5-flare", "size": "1024x1024"},
|
||||
GenericLiteLLMParams(images=references),
|
||||
{},
|
||||
)
|
||||
assert data["model"] == "gpt-image-2"
|
||||
assert data["size"] == "1024x1024"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("as_tuple", [False, True])
|
||||
def test_edit_accepts_filesystem_path(tmp_path, as_tuple):
|
||||
image = tmp_path / "reference.png"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue