mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(replicate): assert same origin on prediction poll URL
urls.get is polled with the operator Replicate token. Reject off-origin values so a poisoned create response cannot become credentialed SSRF.
This commit is contained in:
parent
6a919aec6a
commit
070522dd10
2 changed files with 55 additions and 0 deletions
|
|
@ -7,6 +7,7 @@ from litellm.constants import REPLICATE_MODEL_NAME_WITH_ID_LENGTH
|
|||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
custom_prompt,
|
||||
prompt_factory,
|
||||
|
|
@ -284,6 +285,10 @@ class ReplicateConfig(BaseConfig):
|
|||
...,
|
||||
"urls":{"cancel":"https://api.replicate.com/v1/predictions/gqsmqmp1pdrj00cknr08dgmvb4/cancel","get":"https://api.replicate.com/v1/predictions/gqsmqmp1pdrj00cknr08dgmvb4","stream":"https://stream-b.svc.rno2.c.replicate.net/v1/streams/eot4gbydowuin4snhncydwxt57dfwgsc3w3snycx5nid7oef7jga"}
|
||||
}
|
||||
|
||||
The ``urls.get`` value is polled with the operator's Replicate token.
|
||||
Reject off-origin URLs so a poisoned create response cannot turn the
|
||||
proxy into credentialed SSRF (same class as Azure operation-location).
|
||||
"""
|
||||
response_json: Final = response.json()
|
||||
prediction_url: Final = response_json.get("urls", {}).get("get")
|
||||
|
|
@ -293,6 +298,19 @@ class ReplicateConfig(BaseConfig):
|
|||
message=f"LiteLLM Error - prediction url is None - {response_json}",
|
||||
headers=response.headers,
|
||||
)
|
||||
expected_origin = (
|
||||
str(response.request.url)
|
||||
if response.request is not None
|
||||
else "https://api.replicate.com"
|
||||
)
|
||||
try:
|
||||
assert_same_origin(prediction_url, expected_origin)
|
||||
except SSRFError as ssrf_err:
|
||||
raise ReplicateError(
|
||||
status_code=502,
|
||||
message=f"Rejected prediction URL: {ssrf_err}",
|
||||
headers=response.headers,
|
||||
)
|
||||
return prediction_url
|
||||
|
||||
def validate_environment(
|
||||
|
|
|
|||
|
|
@ -0,0 +1,37 @@
|
|||
"""Tests for Replicate get_prediction_url same-origin guard."""
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.replicate.chat.transformation import ReplicateConfig
|
||||
from litellm.llms.replicate.common_utils import ReplicateError
|
||||
|
||||
|
||||
def _response(urls_get: str | None, request_url: str = "https://api.replicate.com/v1/models/x/predictions"):
|
||||
body = {"urls": {"get": urls_get}} if urls_get is not None else {"urls": {}}
|
||||
request = httpx.Request("POST", request_url)
|
||||
return httpx.Response(200, json=body, request=request)
|
||||
|
||||
|
||||
class TestReplicatePredictionUrlSameOrigin:
|
||||
def test_accepts_same_origin_get_url(self):
|
||||
url = ReplicateConfig().get_prediction_url(
|
||||
_response("https://api.replicate.com/v1/predictions/abc123")
|
||||
)
|
||||
assert url == "https://api.replicate.com/v1/predictions/abc123"
|
||||
|
||||
def test_rejects_off_origin_get_url(self):
|
||||
with pytest.raises(ReplicateError, match="Rejected prediction URL"):
|
||||
ReplicateConfig().get_prediction_url(
|
||||
_response("https://evil.example/steal-token")
|
||||
)
|
||||
|
||||
def test_rejects_http_scheme_mismatch(self):
|
||||
with pytest.raises(ReplicateError, match="Rejected prediction URL"):
|
||||
ReplicateConfig().get_prediction_url(
|
||||
_response("http://api.replicate.com/v1/predictions/abc123")
|
||||
)
|
||||
|
||||
def test_missing_get_url_still_400(self):
|
||||
with pytest.raises(ReplicateError, match="prediction url is None"):
|
||||
ReplicateConfig().get_prediction_url(_response(None))
|
||||
Loading…
Add table
Reference in a new issue