mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
* fix(lens): preserve full trace access and expose investigation failures * chore(lens): sync worker registration schema * fix(lens): support durations without a configured maximum * fix(lens): expose every page of fetched investigation evidence * fix(lens): preserve repeated trace content and interrupt cancelled runs * fix(lens): retry transient heartbeat failures during analysis
419 lines
19 KiB
Python
419 lines
19 KiB
Python
import asyncio
|
|
from queue import SimpleQueue
|
|
from typing import Final
|
|
|
|
import httpx
|
|
import pytest
|
|
from pydantic import ValidationError
|
|
|
|
from litellm.proxy.lens.models import (
|
|
Claim,
|
|
Execution,
|
|
ExecutionContent,
|
|
ModelRequest,
|
|
ModelResult,
|
|
Result,
|
|
Sample,
|
|
TracePart,
|
|
)
|
|
from litellm.proxy.lens.state import queue_job
|
|
from litellm.proxy.lens.worker import LensWorker, failure_message
|
|
from tests.unit.proxy.lens.test_state import NOW, lens
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("failure", (429, 502, 503, 504, "timeout", 402, 409, 401))
|
|
async def test_model_retries_transient_failures_but_not_budget_or_revocation(failure: int | str) -> None:
|
|
attempts: Final = SimpleQueue[str]()
|
|
delays: Final = SimpleQueue[float]()
|
|
expected: Final = ModelResult(content='{"observations":[]}', cost=0.01)
|
|
|
|
def handle(request: httpx.Request) -> httpx.Response:
|
|
attempts.put(request.url.path)
|
|
if attempts.qsize() == 1:
|
|
if failure == "timeout":
|
|
raise httpx.ReadTimeout("upstream timeout", request=request)
|
|
assert isinstance(failure, int)
|
|
return httpx.Response(failure)
|
|
return httpx.Response(200, json=expected.model_dump())
|
|
|
|
async def sleep(delay: float) -> None:
|
|
delays.put(delay)
|
|
|
|
async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client:
|
|
worker: Final = LensWorker(client, sleep=sleep)
|
|
if failure in (402, 409, 401):
|
|
with pytest.raises(httpx.HTTPStatusError):
|
|
await worker.model_request("/model", ModelRequest(purpose="extract", prompt="review"))
|
|
assert attempts.qsize() == 1 and delays.empty()
|
|
else:
|
|
assert await worker.model_request("/model", ModelRequest(purpose="extract", prompt="review")) == expected
|
|
assert attempts.qsize() == 2
|
|
assert delays.get_nowait() == 1 and delays.empty()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_transient_retries_are_bounded() -> None:
|
|
attempts: Final = SimpleQueue[str]()
|
|
delays: Final = SimpleQueue[float]()
|
|
|
|
def handle(request: httpx.Request) -> httpx.Response:
|
|
attempts.put(request.url.path)
|
|
return httpx.Response(503)
|
|
|
|
async def sleep(delay: float) -> None:
|
|
delays.put(delay)
|
|
|
|
async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client:
|
|
with pytest.raises(httpx.HTTPStatusError):
|
|
await LensWorker(client, sleep=sleep).model_request(
|
|
"/model", ModelRequest(purpose="extract", prompt="review")
|
|
)
|
|
assert attempts.qsize() == 3
|
|
assert tuple(delays.get_nowait() for _ in range(delays.qsize())) == (1, 2)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_idle_worker_does_not_start_an_analysis() -> None:
|
|
def handle(request: httpx.Request) -> httpx.Response:
|
|
assert request.url.path == "/lens/worker/claim"
|
|
return httpx.Response(200, content="null")
|
|
|
|
async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client:
|
|
assert await LensWorker(client).run_once() is False
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("result_status", (200, 409))
|
|
async def test_incompatible_claim_reports_failure_instead_of_leaving_the_investigation_running(
|
|
result_status: int,
|
|
) -> None:
|
|
claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=())
|
|
payload: Final = claim.model_dump(mode="json") | {
|
|
"job": claim.job.model_dump(mode="json")
|
|
| {
|
|
"settings": claim.job.settings.model_dump() | {"future_setting": "private content"},
|
|
},
|
|
}
|
|
saved: Final = SimpleQueue[Result]()
|
|
|
|
def handle(request: httpx.Request) -> httpx.Response:
|
|
if request.url.path == "/lens/worker/claim":
|
|
return httpx.Response(200, json=payload)
|
|
assert request.url.path == "/lens/worker/lens/job/result"
|
|
saved.put(Result.model_validate_json(request.content))
|
|
return httpx.Response(result_status, json=True)
|
|
|
|
async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client:
|
|
assert await LensWorker(client).run_once() is True
|
|
assert saved.get_nowait().error == (
|
|
"The worker could not read this investigation. Update the worker to match the gateway, then retry."
|
|
)
|
|
assert saved.empty()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_claim_without_an_identity_does_not_report_failure_for_another_investigation() -> None:
|
|
def handle(request: httpx.Request) -> httpx.Response:
|
|
assert request.url.path == "/lens/worker/claim"
|
|
return httpx.Response(200, json={"job": {"settings": {"future_setting": True}}})
|
|
|
|
async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client:
|
|
with pytest.raises(ValidationError):
|
|
await LensWorker(client).run_once()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("model_status", (200, 402, 503))
|
|
async def test_worker_reads_claimed_activity_and_reports_analysis_or_failure(model_status: int) -> None:
|
|
claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=())
|
|
execution: Final = Execution(
|
|
id="run", source="traces", trace_id="trace", team_id="alpha", name="review", start_time="", span_count=1
|
|
)
|
|
sample: Final = Sample(executions=(execution,), eligible=1)
|
|
content: Final = ExecutionContent(
|
|
execution=execution,
|
|
parts=(TracePart(execution_id="run", span_id="span", name="lead", kind="agent", content="Completed"),),
|
|
)
|
|
saved: Final = SimpleQueue[Result]()
|
|
|
|
def handle(request: httpx.Request) -> httpx.Response:
|
|
match request.url.path:
|
|
case "/lens/worker/claim":
|
|
return httpx.Response(200, json=claim.model_dump(mode="json"))
|
|
case "/lens/worker/lens/job/sample":
|
|
return httpx.Response(200, json=sample.model_dump(mode="json"))
|
|
case "/lens/worker/lens/job/content":
|
|
assert request.url.params["execution_id"] == execution.id
|
|
return httpx.Response(200, json=content.model_dump(mode="json"))
|
|
case "/lens/worker/lens/job/model":
|
|
return httpx.Response(
|
|
model_status,
|
|
json=ModelResult(content='{"observations":[],"cannot_assess":false}', cost=0.01).model_dump(),
|
|
)
|
|
case "/lens/worker/lens/job/progress":
|
|
return httpx.Response(200, json=True)
|
|
case "/lens/worker/lens/job/result":
|
|
saved.put(Result.model_validate_json(request.content))
|
|
return httpx.Response(200, json=True)
|
|
case _:
|
|
pytest.fail(f"Unexpected analyzer request: {request.url.path}")
|
|
|
|
async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client:
|
|
assert await LensWorker(client).run_once() is True
|
|
result: Final = saved.get_nowait()
|
|
assert saved.empty()
|
|
if model_status == 200:
|
|
assert result.error == ""
|
|
assert result.coverage.screened == 1
|
|
assert result.coverage.unassessable == 0
|
|
elif model_status == 402:
|
|
assert "HTTP 402" in result.error and "remaining budget" in result.error
|
|
else:
|
|
assert result.error.startswith("Model request failed (HTTP 503).")
|
|
|
|
|
|
@pytest.mark.parametrize("status", (400, 401, 402, 403, 404, 409, 429, 503))
|
|
def test_failure_reports_action_and_status_without_private_response_content(status: int) -> None:
|
|
request: Final = httpx.Request(
|
|
"POST", "https://private-host.test/lens/worker/private-lens/private-run/model?token=secret"
|
|
)
|
|
response: Final = httpx.Response(status, request=request, text="private trace content and key")
|
|
error: Final = httpx.HTTPStatusError("private exception details", request=request, response=response)
|
|
message: Final = failure_message(error)
|
|
assert message.startswith(f"Model request failed (HTTP {status}).")
|
|
assert "private" not in message and "secret" not in message
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"route,action", (("sample", "Reading trace data"), ("content", "Reading trace data"), ("result", "Saving results"))
|
|
)
|
|
def test_failure_identifies_the_failing_worker_operation(route: str, action: str) -> None:
|
|
request: Final = httpx.Request("GET", f"https://proxy.test/lens/worker/lens/job/{route}")
|
|
response: Final = httpx.Response(503, request=request)
|
|
error: Final = httpx.HTTPStatusError("private body", request=request, response=response)
|
|
assert failure_message(error).startswith(f"{action} failed (HTTP 503).")
|
|
|
|
|
|
def test_connection_timeout_and_invalid_response_have_distinct_private_diagnostics() -> None:
|
|
assert "connect to the proxy" in failure_message(httpx.ConnectError("private hostname"))
|
|
assert "timed out" in failure_message(httpx.ReadTimeout("private prompt"))
|
|
assert "structured JSON" in failure_message(ValueError("private model response"))
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"purpose,stage,schema",
|
|
(
|
|
("extract", "Reading executions", "TraceReview"),
|
|
("cluster", "Grouping observations", "Clusters"),
|
|
("investigate", "Checking original evidence", "Decision"),
|
|
),
|
|
)
|
|
async def test_worker_saves_validation_errors_from_every_analysis_stage(purpose: str, stage: str, schema: str) -> None:
|
|
import json
|
|
|
|
claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=())
|
|
execution: Final = Execution(
|
|
id="run", source="traces", trace_id="trace", team_id="alpha", name="review", start_time="", span_count=1
|
|
)
|
|
sample: Final = Sample(executions=(execution,), eligible=1)
|
|
content: Final = ExecutionContent(
|
|
execution=execution,
|
|
parts=(TracePart(execution_id="run", span_id="span", name="lead", kind="agent", content="Tool timeout"),),
|
|
)
|
|
saved: Final = SimpleQueue[Result]()
|
|
attempts: Final = SimpleQueue[str]()
|
|
|
|
def handle(request: httpx.Request) -> httpx.Response:
|
|
match request.url.path.rsplit("/", 1)[-1]:
|
|
case "claim":
|
|
return httpx.Response(200, json=claim.model_dump(mode="json"))
|
|
case "sample":
|
|
return httpx.Response(200, json=sample.model_dump(mode="json"))
|
|
case "content":
|
|
return httpx.Response(200, json=content.model_dump(mode="json"))
|
|
case "model":
|
|
body: Final = ModelRequest.model_validate_json(request.content)
|
|
if body.purpose == purpose:
|
|
attempts.put(body.purpose)
|
|
return httpx.Response(
|
|
200,
|
|
json={"content": '{"candidates":[', "cost": 0.01},
|
|
headers={"x-litellm-lens-finish-reason": "length"},
|
|
)
|
|
if body.purpose == "cluster":
|
|
return httpx.Response(
|
|
200,
|
|
json={
|
|
"content": json.dumps({"candidates": json.loads(body.prompt)["candidates"]}),
|
|
"cost": 0.01,
|
|
},
|
|
)
|
|
return httpx.Response(
|
|
200,
|
|
json={
|
|
"content": json.dumps(
|
|
{
|
|
"observations": [
|
|
{
|
|
"check_id": claim.job.settings.analysis_checks[0].id,
|
|
"summary": "Tool timeout",
|
|
"evidence": [
|
|
{"execution_id": "r0", "span_id": "span", "quote": "Tool timeout"}
|
|
],
|
|
}
|
|
]
|
|
}
|
|
),
|
|
"cost": 0.01,
|
|
},
|
|
)
|
|
case "progress":
|
|
return httpx.Response(200, json=True)
|
|
case "result":
|
|
saved.put(Result.model_validate_json(request.content))
|
|
return httpx.Response(200, json=True)
|
|
case _:
|
|
pytest.fail(f"Unexpected worker request: {request.url.path}")
|
|
|
|
async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client:
|
|
assert await LensWorker(client).run_once()
|
|
message: Final = saved.get_nowait().error
|
|
assert message.startswith(f"{stage} failed: {schema} response invalid after 2 attempts.")
|
|
assert "finish_reason=length" in message
|
|
assert "EOF while parsing" in message and "[json_invalid]" in message
|
|
assert attempts.qsize() == 2 and saved.empty()
|
|
|
|
|
|
def test_response_validation_diagnostics_omit_input_values_and_unexpected_field_names() -> None:
|
|
with pytest.raises(ValidationError) as caught:
|
|
ModelResult.model_validate({"content": "private trace", "cost": "private token", "private field": "secret"})
|
|
message: Final = failure_message(caught.value)
|
|
assert "Invalid ModelResult response" in message
|
|
assert "cost:" in message and "[float_parsing]" in message
|
|
assert "[extra_forbidden]" in message
|
|
assert "private" not in message and "secret" not in message
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("heartbeat_status", (401, 403, 409))
|
|
async def test_losing_the_lease_interrupts_an_in_flight_model_request(heartbeat_status: int) -> None:
|
|
claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=())
|
|
execution: Final = Execution(
|
|
id="run", source="traces", trace_id="t", team_id="", name="task", start_time="", span_count=1
|
|
)
|
|
started: Final = asyncio.Event()
|
|
cancelled: Final = asyncio.Event()
|
|
never: Final = asyncio.Event()
|
|
saved: Final = SimpleQueue[Result]()
|
|
|
|
async def heartbeat_wait(_seconds: float) -> None:
|
|
await started.wait()
|
|
|
|
async def handle(request: httpx.Request) -> httpx.Response:
|
|
match request.url.path.rsplit("/", 1)[-1]:
|
|
case "claim":
|
|
return httpx.Response(200, json=claim.model_dump(mode="json"))
|
|
case "sample":
|
|
return httpx.Response(200, json=Sample(executions=(execution,), eligible=1).model_dump())
|
|
case "content":
|
|
return httpx.Response(
|
|
200,
|
|
json=ExecutionContent(
|
|
execution=execution,
|
|
parts=(
|
|
TracePart(execution_id="run", span_id="span", name="step", kind="tool", content="evidence"),
|
|
),
|
|
).model_dump(),
|
|
)
|
|
case "model":
|
|
assert request.extensions["timeout"] == {"connect": 13, "read": None, "write": 13, "pool": 13}
|
|
started.set()
|
|
try:
|
|
await never.wait()
|
|
finally:
|
|
cancelled.set()
|
|
pytest.fail("The cancelled model request must not finish")
|
|
case "heartbeat":
|
|
return httpx.Response(heartbeat_status)
|
|
case "progress":
|
|
return httpx.Response(200, json=True)
|
|
case "result":
|
|
saved.put(Result.model_validate_json(request.content))
|
|
return httpx.Response(409)
|
|
case _:
|
|
pytest.fail(f"Unexpected worker request: {request.url.path}")
|
|
|
|
async with httpx.AsyncClient(
|
|
base_url="https://proxy.test", transport=httpx.MockTransport(handle), timeout=13
|
|
) as client:
|
|
assert await LensWorker(client, heartbeat_wait=heartbeat_wait).run_once()
|
|
assert cancelled.is_set()
|
|
assert f"HTTP {heartbeat_status}" in saved.get_nowait().error
|
|
assert saved.empty()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("failure", (429, 500, 502, 503, 504, "connection", "timeout"))
|
|
async def test_transient_heartbeat_failure_recovers_without_cancelling_analysis(failure: int | str) -> None:
|
|
claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=())
|
|
execution: Final = Execution(
|
|
id="run", source="traces", trace_id="t", team_id="", name="task", start_time="", span_count=1
|
|
)
|
|
started: Final = asyncio.Event()
|
|
recovered: Final = asyncio.Event()
|
|
never: Final = asyncio.Event()
|
|
attempts: Final = SimpleQueue[str]()
|
|
saved: Final = SimpleQueue[Result]()
|
|
|
|
async def heartbeat_wait(_seconds: float) -> None:
|
|
await started.wait()
|
|
if attempts.qsize() >= 2:
|
|
await never.wait()
|
|
|
|
async def handle(request: httpx.Request) -> httpx.Response:
|
|
match request.url.path.rsplit("/", 1)[-1]:
|
|
case "claim":
|
|
return httpx.Response(200, json=claim.model_dump(mode="json"))
|
|
case "sample":
|
|
return httpx.Response(200, json=Sample(executions=(execution,), eligible=1).model_dump())
|
|
case "content":
|
|
return httpx.Response(
|
|
200,
|
|
json=ExecutionContent(
|
|
execution=execution,
|
|
parts=(
|
|
TracePart(execution_id="run", span_id="span", name="step", kind="tool", content="evidence"),
|
|
),
|
|
).model_dump(),
|
|
)
|
|
case "model":
|
|
started.set()
|
|
await recovered.wait()
|
|
return httpx.Response(200, json={"content": '{"observations":[],"cannot_assess":false}', "cost": 0.01})
|
|
case "heartbeat":
|
|
attempts.put(request.url.path)
|
|
if attempts.qsize() == 1:
|
|
if failure == "connection":
|
|
raise httpx.ConnectError("temporary connection failure", request=request)
|
|
if failure == "timeout":
|
|
raise httpx.ReadTimeout("temporary response timeout", request=request)
|
|
assert isinstance(failure, int)
|
|
return httpx.Response(failure)
|
|
recovered.set()
|
|
return httpx.Response(200, json=True)
|
|
case "progress":
|
|
return httpx.Response(200, json=True)
|
|
case "result":
|
|
saved.put(Result.model_validate_json(request.content))
|
|
return httpx.Response(200, json=True)
|
|
case _:
|
|
pytest.fail(f"Unexpected worker request: {request.url.path}")
|
|
|
|
async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client:
|
|
assert await LensWorker(client, heartbeat_wait=heartbeat_wait).run_once()
|
|
result: Final = saved.get_nowait()
|
|
assert result.error == ""
|
|
assert result.coverage.screened == 1 and result.coverage.unassessable == 0
|
|
assert attempts.qsize() == 2 and saved.empty()
|