mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
Merge ae750ad7a6 into 6a3333d3c8
This commit is contained in:
commit
ce99fe71f9
2 changed files with 84 additions and 2 deletions
|
|
@ -19,6 +19,7 @@ Admins can opt out via two ``litellm`` globals (wired from proxy config):
|
|||
check but still resolve DNS and still rewrite HTTP to the resolved IP.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import socket
|
||||
from ipaddress import ip_address, ip_network
|
||||
from typing import Any, Final
|
||||
|
|
@ -413,14 +414,21 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any:
|
|||
|
||||
|
||||
async def async_safe_get(client: Any, url: str, **kwargs: Any) -> Any:
|
||||
"""Async version of safe_get."""
|
||||
"""Async version of safe_get.
|
||||
|
||||
``validate_url`` resolves DNS through the blocking ``socket.getaddrinfo``,
|
||||
and the hostname comes from the caller-supplied URL, so run it in a worker
|
||||
thread. Resolving on the event loop lets one URL pointed at a stalling
|
||||
nameserver pin the loop for the resolver timeout and stall every other
|
||||
in-flight request on the worker.
|
||||
"""
|
||||
if not getattr(litellm, "user_url_validation", True):
|
||||
kwargs.setdefault("follow_redirects", True)
|
||||
return await client.get(url, **kwargs)
|
||||
kwargs.pop("follow_redirects", None)
|
||||
caller_headers: Final = kwargs.pop("headers", {})
|
||||
for _ in range(_MAX_REDIRECTS):
|
||||
validated_url, original_host = validate_url(url)
|
||||
validated_url, original_host = await asyncio.to_thread(validate_url, url)
|
||||
response = await client.get(
|
||||
validated_url,
|
||||
headers={**caller_headers, "Host": original_host},
|
||||
|
|
|
|||
|
|
@ -535,3 +535,77 @@ def test_assert_same_origin_error_message_does_not_leak_hostnames():
|
|||
detail = str(exc.value)
|
||||
assert "attacker.example.com" not in detail
|
||||
assert "api.internal-corp.example" not in detail
|
||||
|
||||
|
||||
class TestAsyncSafeGetDoesNotBlockEventLoop:
|
||||
"""``validate_url`` calls the blocking ``socket.getaddrinfo``.
|
||||
|
||||
Every ``async_safe_get`` caller fetches a URL that came from the request
|
||||
body (``image_url``/``file_url`` on a chat completion, an A2A agent-card
|
||||
URL, an OpenAPI spec URL for an MCP server, a RAG ingestion file URL), so
|
||||
an unprivileged caller picks the hostname that gets resolved. If that
|
||||
resolution runs on the event loop, one caller pointing at a stalling
|
||||
authoritative nameserver freezes every other in-flight request on the
|
||||
worker for the resolver timeout.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _fake_response():
|
||||
class _Resp:
|
||||
is_redirect = False
|
||||
status_code = 200
|
||||
|
||||
return _Resp()
|
||||
|
||||
class _FakeAsyncClient:
|
||||
async def get(self, url, **kwargs):
|
||||
return TestAsyncSafeGetDoesNotBlockEventLoop._fake_response()
|
||||
|
||||
async def test_event_loop_stays_responsive_during_dns_resolution(
|
||||
self, monkeypatch
|
||||
):
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
dns_delay = 0.5
|
||||
|
||||
def slow_getaddrinfo(host, port, *args, **kwargs):
|
||||
time.sleep(dns_delay)
|
||||
return [
|
||||
(
|
||||
socket.AF_INET,
|
||||
socket.SOCK_STREAM,
|
||||
6,
|
||||
"",
|
||||
("93.184.216.34", port or 443),
|
||||
)
|
||||
]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", slow_getaddrinfo)
|
||||
|
||||
ticks = 0
|
||||
stop = False
|
||||
|
||||
async def heartbeat():
|
||||
nonlocal ticks
|
||||
while not stop:
|
||||
await asyncio.sleep(0.01)
|
||||
ticks += 1
|
||||
|
||||
beat = asyncio.create_task(heartbeat())
|
||||
try:
|
||||
response = await url_utils.async_safe_get(
|
||||
self._FakeAsyncClient(), "https://images.example.com/cat.png"
|
||||
)
|
||||
finally:
|
||||
stop = True
|
||||
await beat
|
||||
|
||||
assert response.status_code == 200
|
||||
# ~50 ticks are due while DNS is in flight; require a fraction of
|
||||
# that so the assertion is not timing-fragile, but 0-1 ticks (the
|
||||
# loop pinned for the whole resolution) still fails.
|
||||
assert ticks >= 10, (
|
||||
f"event loop was blocked during DNS resolution: only {ticks} "
|
||||
"heartbeat ticks fired while validate_url resolved the hostname"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue