fix(lens): reuse control connections and satisfy review checks

This commit is contained in:
moe-berri 2026-10-07 13:33:00 -07:00
parent 23ac6d7af9
commit 78c6237d0c
7 changed files with 74 additions and 12 deletions

View file

@ -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::{

View file

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

View file

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

View file

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

View file

@ -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",

View file

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

View file

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