mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(vertex_ai): route endpoint resolution through safe_get to guard caller-supplied api_base
This commit is contained in:
parent
54fd69beb2
commit
18373f8f51
2 changed files with 30 additions and 11 deletions
|
|
@ -218,7 +218,17 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
model=model,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
response: Final = sync_handler.get(url=endpoint_url, headers=headers)
|
||||
# ``api_base`` can come from caller-supplied request kwargs, so wrap the
|
||||
# fetch in ``safe_get``: it rejects DNS-rebind / private / cloud-metadata
|
||||
# targets before the bearer token leaves the process (mirrors retrieve_batch).
|
||||
fetched: Final[_FetchedResponseView] = {
|
||||
"response": safe_get(
|
||||
sync_handler,
|
||||
endpoint_url,
|
||||
headers=headers,
|
||||
)
|
||||
}
|
||||
response: Final = fetched["response"]
|
||||
if response.status_code != 200:
|
||||
raise VertexAIError(
|
||||
status_code=response.status_code,
|
||||
|
|
|
|||
|
|
@ -184,7 +184,10 @@ def test_create_batch_sync_does_not_resolve_publisher_models():
|
|||
client = MagicMock()
|
||||
client.post.return_value = _http_response()
|
||||
|
||||
with patch(f"{HMOD}._get_httpx_client", return_value=client):
|
||||
with (
|
||||
patch(f"{HMOD}._get_httpx_client", return_value=client),
|
||||
patch(f"{HMOD}.safe_get") as safe_get,
|
||||
):
|
||||
out = h.create_batch(
|
||||
_is_async=False,
|
||||
create_batch_data=CREATE_DATA,
|
||||
|
|
@ -199,7 +202,7 @@ def test_create_batch_sync_does_not_resolve_publisher_models():
|
|||
assert isinstance(out, LiteLLMBatch)
|
||||
sent = json.loads(client.post.call_args.kwargs["data"])
|
||||
assert sent["model"] == "publishers/google/models/gemini-1.5-flash-001"
|
||||
client.get.assert_not_called()
|
||||
safe_get.assert_not_called()
|
||||
|
||||
|
||||
ENDPOINT_ID = "7768560373388541952"
|
||||
|
|
@ -226,10 +229,12 @@ def test_create_batch_sync_resolves_fine_tuned_endpoint_to_tuned_model():
|
|||
tuned model resource; the v1 batch API rejects endpoint resources in `model` (LIT-6899)."""
|
||||
h = _make_handler()
|
||||
client = MagicMock()
|
||||
client.get.return_value = _endpoint_get_response()
|
||||
client.post.return_value = _http_response()
|
||||
|
||||
with patch(f"{HMOD}._get_httpx_client", return_value=client):
|
||||
with (
|
||||
patch(f"{HMOD}._get_httpx_client", return_value=client),
|
||||
patch(f"{HMOD}.safe_get", return_value=_endpoint_get_response()) as safe_get,
|
||||
):
|
||||
out = h.create_batch(
|
||||
_is_async=False,
|
||||
create_batch_data=ENDPOINT_CREATE_DATA,
|
||||
|
|
@ -242,8 +247,8 @@ def test_create_batch_sync_resolves_fine_tuned_endpoint_to_tuned_model():
|
|||
)
|
||||
|
||||
assert isinstance(out, LiteLLMBatch)
|
||||
get_kwargs = client.get.call_args.kwargs
|
||||
assert get_kwargs["url"] == (
|
||||
get_args, get_kwargs = safe_get.call_args
|
||||
assert get_args[1] == (
|
||||
f"https://{LOCATION}-aiplatform.googleapis.com/v1/projects/{PROJECT}"
|
||||
f"/locations/{LOCATION}/endpoints/{ENDPOINT_ID}"
|
||||
)
|
||||
|
|
@ -291,9 +296,11 @@ def test_create_batch_sync_endpoint_resolution_error_raises():
|
|||
resolve_response = MagicMock()
|
||||
resolve_response.status_code = 404
|
||||
resolve_response.text = "endpoint not found"
|
||||
client.get.return_value = resolve_response
|
||||
|
||||
with patch(f"{HMOD}._get_httpx_client", return_value=client):
|
||||
with (
|
||||
patch(f"{HMOD}._get_httpx_client", return_value=client),
|
||||
patch(f"{HMOD}.safe_get", return_value=resolve_response),
|
||||
):
|
||||
with pytest.raises(VertexAIError) as exc_info:
|
||||
h.create_batch(
|
||||
_is_async=False,
|
||||
|
|
@ -339,9 +346,11 @@ def test_create_batch_custom_endpoint_raises_400_without_io():
|
|||
def test_create_batch_sync_endpoint_without_deployed_model_raises_400():
|
||||
h = _make_handler()
|
||||
client = MagicMock()
|
||||
client.get.return_value = _endpoint_get_response(deployed_models=[])
|
||||
|
||||
with patch(f"{HMOD}._get_httpx_client", return_value=client):
|
||||
with (
|
||||
patch(f"{HMOD}._get_httpx_client", return_value=client),
|
||||
patch(f"{HMOD}.safe_get", return_value=_endpoint_get_response(deployed_models=[])),
|
||||
):
|
||||
with pytest.raises(VertexAIError) as exc_info:
|
||||
h.create_batch(
|
||||
_is_async=False,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue