mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(lens): reuse control connections and satisfy review checks
This commit is contained in:
parent
23ac6d7af9
commit
78c6237d0c
7 changed files with 74 additions and 12 deletions
|
|
@ -287,6 +287,7 @@ async fn configured_private_dns_names_are_reachable_without_following_redirects(
|
|||
assert_eq!(redirected.status(), 302);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn oversized_combined_tool_replies_remain_readable_after_a_checkpoint() {
|
||||
use litellm_lens::{
|
||||
|
|
|
|||
|
|
@ -194,10 +194,10 @@ async def service_connection(auth: Auth) -> ServiceConnection:
|
|||
public_url: Final = os.environ.get("LITELLM_LENS_PUBLIC_URL", "").rstrip("/")
|
||||
try:
|
||||
connection: Final = LensConnection.from_env()
|
||||
async with (
|
||||
connection.client() as client,
|
||||
client.stream("GET", "/internal/status", timeout=2) as response,
|
||||
):
|
||||
client: Final = connection.control_client()
|
||||
async with client.stream(
|
||||
"GET", connection.endpoint("/internal/status"), headers=connection.headers, timeout=2
|
||||
) as response:
|
||||
if response.status_code == 200:
|
||||
status: Final = ServiceStatus.model_validate_json(await bounded_response(response, 16 * 1024))
|
||||
return ServiceConnection(url=public_url, connected=True, status=status)
|
||||
|
|
@ -225,11 +225,13 @@ async def publish_credentials() -> bool:
|
|||
try:
|
||||
connection: Final = LensConnection.from_env()
|
||||
snapshot: Final = await credential_snapshot()
|
||||
async with connection.client() as client:
|
||||
response: Final = await client.post(
|
||||
"/internal/credentials", json=snapshot.model_dump(mode="json"), timeout=2
|
||||
)
|
||||
return response.status_code == 204
|
||||
response: Final = await connection.control_client().post(
|
||||
connection.endpoint("/internal/credentials"),
|
||||
headers=connection.headers,
|
||||
json=snapshot.model_dump(mode="json"),
|
||||
timeout=2,
|
||||
)
|
||||
return response.status_code == 204
|
||||
except (ValueError, httpx.HTTPError):
|
||||
return False
|
||||
|
||||
|
|
|
|||
|
|
@ -39,7 +39,7 @@ async def manage_tracing(
|
|||
enabled: bool,
|
||||
receiver_factory: Callable[[], TraceReceiver] | None = None,
|
||||
settings: Mapping[str, object] | None = None,
|
||||
client_factory: Callable[[LensConnection], httpx.AsyncClient] = LensConnection.client,
|
||||
client_factory: Callable[[LensConnection], httpx.AsyncClient] = LensConnection.lifespan_client,
|
||||
) -> AsyncGenerator[TraceReceiver | None, None]:
|
||||
if not enabled:
|
||||
yield None
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import httpx
|
|||
from pydantic import JsonValue, TypeAdapter
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.rust_bridge.trace.errors import TraceChanged
|
||||
from litellm.rust_bridge.trace.generated.types import QueryScope, ReadQueryName, TraceScope
|
||||
|
||||
|
|
@ -39,10 +40,23 @@ class LensConnection:
|
|||
raise ValueError("Set LITELLM_LENS_SERVICE_TOKEN to the same secret on LiteLLM and Lens")
|
||||
return cls(url, token)
|
||||
|
||||
def client(self) -> httpx.AsyncClient:
|
||||
def control_client(self) -> httpx.AsyncClient:
|
||||
return get_async_httpx_client(
|
||||
"lens-control",
|
||||
params={"timeout": httpx.Timeout(35, connect=3), "follow_redirects": False},
|
||||
).client
|
||||
|
||||
def endpoint(self, path: str) -> str:
|
||||
return self.url + path
|
||||
|
||||
@property
|
||||
def headers(self) -> Mapping[str, str]:
|
||||
return {"Authorization": f"Bearer {self.token}"}
|
||||
|
||||
def lifespan_client(self) -> httpx.AsyncClient:
|
||||
return httpx.AsyncClient(
|
||||
base_url=self.url,
|
||||
headers={"Authorization": f"Bearer {self.token}"},
|
||||
headers=self.headers,
|
||||
timeout=httpx.Timeout(35, connect=3),
|
||||
limits=httpx.Limits(max_connections=10, max_keepalive_connections=10),
|
||||
follow_redirects=False,
|
||||
|
|
|
|||
|
|
@ -2,6 +2,9 @@ import ast
|
|||
import os
|
||||
|
||||
ALLOWED_FILES = [
|
||||
# Lens data traffic owns one pool per app lifespan, isolated from model traffic and closed on shutdown.
|
||||
"../../litellm/tracing/remote.py",
|
||||
"./litellm/tracing/remote.py",
|
||||
# local files
|
||||
"../../litellm/__init__.py",
|
||||
"../../litellm/llms/custom_httpx/http_handler.py",
|
||||
|
|
|
|||
|
|
@ -1031,6 +1031,7 @@ async def test_internal_service_authentication_is_separate_from_gateway_keys(
|
|||
(200, b"x" * 17000, False),
|
||||
),
|
||||
)
|
||||
@pytest.mark.usefixtures("httpx_transport")
|
||||
async def test_service_status_uses_internal_auth_and_only_advertises_the_public_url(
|
||||
monkeypatch: pytest.MonkeyPatch, status: int, content: bytes, connected: bool
|
||||
) -> None:
|
||||
|
|
@ -1078,6 +1079,7 @@ async def test_credential_snapshot_excludes_expired_keys_and_disables_caching(mo
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("accepted", (True, False))
|
||||
@pytest.mark.usefixtures("httpx_transport")
|
||||
async def test_created_ingestion_keys_report_activation_only_after_the_service_acknowledges(
|
||||
monkeypatch: pytest.MonkeyPatch, accepted: bool
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -185,3 +185,43 @@ async def test_request_records_use_the_internal_service_endpoint() -> None:
|
|||
request: Final = requests.get_nowait()
|
||||
assert request.url.path == "/internal/spend"
|
||||
assert json.loads(request.content) == [{"request_id": "r"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_control_requests_reuse_connections_without_retaining_another_service_credential() -> None:
|
||||
requests: Final[asyncio.Queue[tuple[str, bytes]]] = asyncio.Queue()
|
||||
|
||||
async def serve(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
|
||||
try:
|
||||
while True:
|
||||
headers: Final = await reader.readuntil(b"\r\n\r\n")
|
||||
requests.put_nowait((str(writer.get_extra_info("peername")), headers))
|
||||
writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\n{}")
|
||||
await writer.drain()
|
||||
except asyncio.IncompleteReadError:
|
||||
pass
|
||||
finally:
|
||||
writer.close()
|
||||
await writer.wait_closed()
|
||||
|
||||
async with await asyncio.start_server(serve, "127.0.0.1", 0) as server:
|
||||
port: Final = server.sockets[0].getsockname()[1]
|
||||
first: Final = LensConnection(f"http://127.0.0.1:{port}/one", "first-service-token")
|
||||
second: Final = LensConnection(f"http://127.0.0.1:{port}/two", "second-service-token")
|
||||
try:
|
||||
for connection in (first, second):
|
||||
response: Final = await connection.control_client().get(
|
||||
connection.endpoint("/internal/status"), headers=connection.headers
|
||||
)
|
||||
assert response.json() == {}
|
||||
first_peer, first_request = await asyncio.wait_for(requests.get(), 2)
|
||||
second_peer, second_request = await asyncio.wait_for(requests.get(), 2)
|
||||
assert first_peer == second_peer
|
||||
assert b"GET /one/internal/status " in first_request
|
||||
assert b"GET /two/internal/status " in second_request
|
||||
assert b"Bearer first-service-token" in first_request
|
||||
assert b"Bearer second-service-token" not in first_request
|
||||
assert b"Bearer second-service-token" in second_request
|
||||
assert b"Bearer first-service-token" not in second_request
|
||||
finally:
|
||||
await second.control_client().aclose()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue