This commit is contained in:
Sash 2026-09-29 14:18:38 -07:00 • committed by GitHub
commit f073b61aa5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 55 additions and 0 deletions

View file

@ -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,
@ -285,6 +286,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")
@ -294,6 +299,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(

View file

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