feat(lens): report a review with reasoning for each screened trace

This commit is contained in:
Ishaan Jaff 2026-10-03 16:08:59 -07:00
parent 9a7c6fe9f4
commit 19cc5035d9
No known key found for this signature in database
5 changed files with 235 additions and 18 deletions

View file

@ -1,11 +1,13 @@
import asyncio
import json
import time
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable
from contextlib import aclosing
from datetime import datetime, timezone
from functools import reduce
from itertools import chain, islice
from types import MappingProxyType
from typing import Final, Literal, TypeAlias, TypeVar
from typing import Final, Literal, Protocol, TypeAlias, TypeVar
from pydantic import Field, TypeAdapter, ValidationError
@ -20,6 +22,9 @@ from .models import (
ModelResult,
Record,
Result,
Review,
ReviewSpan,
ReviewVerdict,
RunAssessment,
Sample,
TracePart,
@ -38,6 +43,7 @@ class Observation(Record):
class Extraction(Record):
observations: tuple[Observation, ...] = ()
cannot_assess: bool = False
reasoning: str = Field(default="", max_length=800)
class SpanRead(Record):
@ -84,6 +90,8 @@ class Examined(Record):
partial: bool
cannot_assess: bool
error: str = ""
reasoning: str = ""
shown: tuple[TracePart, ...] = ()
class Investigation(Record):
@ -94,7 +102,10 @@ class Investigation(Record):
ModelCall: TypeAlias = Callable[[ModelRequest], Awaitable[ModelResult]]
ReadContent: TypeAlias = Callable[[str, str, int], Awaitable[ExecutionContent]]
ReportProgress: TypeAlias = Callable[[str, Coverage], Awaitable[None]]
class ReportProgress(Protocol):
def __call__(self, stage: str, coverage: Coverage, review: Review | None = None, /) -> Awaitable[None]: ...
ResponseT = TypeVar("ResponseT", bound=Record)
@ -326,7 +337,9 @@ async def extract_stored(
request: Final = ModelRequest(purpose="extract", prompt=prompt)
if must_decide:
final: Final = await structured_response(request, Extraction, model)
return TraceReview(observations=final.observations, cannot_assess=final.cannot_assess)
return TraceReview(
observations=final.observations, cannot_assess=final.cannot_assess, reasoning=final.reasoning
)
return await structured_response(request, TraceReview, model)
response: TraceReview
@ -375,6 +388,7 @@ async def extract_stored(
parts=evidence,
partial=page.partial or page.next_cursor is not None or bool(response.reads) or invalid_observations,
cannot_assess=not span_count or response.cannot_assess or bool(response.reads) or invalid_observations,
reasoning=response.reasoning,
)
reviews: Final = tuple([await examine(catalog) for catalog in store.catalogs(root_count)])
@ -383,12 +397,52 @@ async def extract_stored(
retained: Final = tuple(
p for p in chain.from_iterable(r.parts for r in reviews) if p.span_id in cited or not p.parent_span_id
)
leading: Final = MappingProxyType(
{
p.span_id: p
for p in (*((first_root,) if first_root else ()), *(p for p in store.parts() if p.span_id in cited))
}
)
shown: Final = islice(chain(leading.values(), (p for p in store.parts() if p.span_id not in leading)), 8)
return Examined(
execution=execution,
observations=observations,
parts=tuple(dict.fromkeys((*retained, *((first_root,) if first_root else ())))),
partial=any(r.partial for r in reviews),
cannot_assess=not reviews or all(r.cannot_assess for r in reviews),
reasoning=" ".join(r.reasoning for r in reviews if r.reasoning),
shown=tuple(p.model_copy(update=MappingProxyType({"content": overview_content(p, root_count)})) for p in shown),
)
def review_of(examined: Examined, model: str, duration_ms: int, at: datetime) -> Review:
execution: Final = examined.execution
cited: Final = frozenset(
(e.execution_id, e.span_id) for e in chain.from_iterable(o.evidence for o in examined.observations)
)
return Review(
execution_id=execution.id,
trace_id=execution.trace_id,
agent=execution.service or execution.name,
name=execution.name,
spans=tuple(
ReviewSpan(
span_id=p.span_id,
name=p.name[:120],
kind=p.kind[:40],
preview=p.content[:240],
cited=(p.execution_id, p.span_id) in cited,
)
for p in examined.shown[:8]
),
reasoning=examined.reasoning[:800],
verdicts=tuple(
ReviewVerdict(check_id=o.check_id, kind=o.kind, summary=o.summary[:300]) for o in examined.observations
),
cannot_assess=examined.cannot_assess,
model=model,
duration_ms=max(duration_ms, 0),
at=at,
)
@ -630,8 +684,19 @@ async def analyze_sample(
)
)
async def progress_original(stage: str, coverage: Coverage, review: Review | None = None, /) -> None:
await progress(
stage,
coverage,
review and review.model_copy(update=MappingProxyType({"execution_id": originals[review.execution_id].id})),
)
result: Final = await _analyze_sample(
claim, sample.model_copy(update=MappingProxyType({"executions": executions})), read_alias, model, progress
claim,
sample.model_copy(update=MappingProxyType({"executions": executions})),
read_alias,
model,
progress_original,
)
return result.model_copy(
update=MappingProxyType(
@ -851,16 +916,20 @@ async def merge_candidates(
async def examine_executions(
claim: Claim, sample: Sample, read: ReadContent, model: ModelCall, progress: ReportProgress
) -> AsyncIterator[Examined]:
async def examine(execution: Execution) -> Examined:
return await extract(claim, execution, read, model)
async def examine(execution: Execution) -> tuple[Examined, Review]:
started: Final = time.perf_counter()
examined: Final = await extract(claim, execution, read, model)
elapsed: Final = round((time.perf_counter() - started) * 1000)
return examined, review_of(examined, claim.job.settings.model, elapsed, datetime.now(timezone.utc))
await progress("Reading executions", Coverage(eligible=sample.eligible, selected=len(sample.executions)))
completed: Final = iter(range(1, len(sample.executions) + 1))
async with aclosing(concurrent_results(sample.executions, examine, claim.job.settings.concurrency)) as results:
async for item in results:
async for item, review in results:
await progress(
"Reading executions",
Coverage(eligible=sample.eligible, selected=len(sample.executions), screened=next(completed)),
review,
)
yield item

View file

@ -26,3 +26,5 @@ If you need more evidence, return reads; otherwise return reads=[] and your fina
Carry forward still-valid earlier observations and remove disproved ones.
cannot_assess means insufficient evidence to assess this run, not absence of an issue.
Never manufacture an issue just to produce a result.
Set reasoning to 1-3 plain sentences: what the agent was asked, what happened, and why your observations follow, or why the run is fine.
Keep reasoning under 800 characters and do not quote secrets or long trace text in it.

View file

@ -10,7 +10,7 @@ import httpx
from pydantic import BaseModel, ConfigDict, ValidationError
from .analysis import AnalysisResponseError, analyze_sample, validation_details
from .models import Claim, Coverage, ExecutionContent, ModelRequest, ModelResult, Progress, Result, Sample
from .models import Claim, Coverage, ExecutionContent, ModelRequest, ModelResult, Progress, Result, Review, Sample
logger: Final = logging.getLogger("litellm.lens.worker")
@ -160,9 +160,10 @@ class LensWorker:
result.raise_for_status()
return ExecutionContent.model_validate(result.json())
async def progress(stage: str, coverage: Coverage) -> None:
async def progress(stage: str, coverage: Coverage, review: Review | None = None, /) -> None:
result: Final = await self.client.post(
prefix + "/progress", json=Progress(stage=stage, coverage=coverage).model_dump()
prefix + "/progress",
json=Progress(stage=stage, coverage=coverage, review=review).model_dump(mode="json"),
)
result.raise_for_status()

View file

@ -15,6 +15,7 @@ from litellm.proxy.lens.models import (
ExecutionContent,
ModelRequest,
ModelResult,
Review,
Sample,
TracePart,
)
@ -66,7 +67,7 @@ async def test_parallel_review_shares_one_model_limit_and_cleans_up(outcome: str
finally:
exited.put(request.prompt)
async def progress(stage: str, coverage: Coverage) -> None:
async def progress(stage: str, coverage: Coverage, _review: Review | None = None, /) -> None:
if stage == "Reading executions":
counts.put(coverage.screened)
@ -117,7 +118,7 @@ async def test_independent_investigations_overlap_and_report_completions() -> No
async def read(_execution_id: str, _cursor: str, _offset: int) -> ExecutionContent:
pytest.fail("Inconclusive decisions must not fetch evidence")
async def progress(stage: str, coverage: Coverage) -> None:
async def progress(stage: str, coverage: Coverage, _review: Review | None = None, /) -> None:
assert stage == "Checking original evidence"
progress_counts.put(coverage.investigated)
@ -464,7 +465,7 @@ async def test_grouping_consolidates_prior_batches_and_reports_real_progress() -
)
stages: Final = iter((0, 1))
async def progress(stage: str, coverage: Coverage) -> None:
async def progress(stage: str, coverage: Coverage, _review: Review | None = None, /) -> None:
assert stage == "Grouping observations"
assert coverage.grouping_batches == 2
assert coverage.grouped_batches == next(stages)
@ -558,7 +559,7 @@ async def test_thousands_of_matching_runs_keep_all_members_without_a_growing_mod
cost=0,
)
async def progress(_stage: str, coverage: Coverage) -> None:
async def progress(_stage: str, coverage: Coverage, _review: Review | None = None, /) -> None:
counts.put(coverage.grouped_batches)
batches: Final = observation_batches(observations)
@ -632,7 +633,7 @@ async def test_review_keeps_original_ids_in_per_run_assessments() -> None:
async def model(_request: ModelRequest) -> ModelResult:
return ModelResult(content='{"observations":[],"cannot_assess":false}', cost=0)
async def progress(_stage: str, _coverage: Coverage) -> None:
async def progress(_stage: str, _coverage: Coverage, _review: Review | None = None, /) -> None:
pass
claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=())
@ -891,7 +892,7 @@ async def test_final_registry_reconciles_patterns_split_across_pages() -> None:
)
return ModelResult(content=Clusters(candidates=grouped).model_dump_json(), cost=0)
async def progress(_stage: str, _coverage: Coverage) -> None:
async def progress(_stage: str, _coverage: Coverage, _review: Review | None = None, /) -> None:
return None
result: Final = await cluster_batches((observations,), model, progress, Coverage())
@ -918,7 +919,7 @@ async def test_distinct_patterns_are_consolidated_in_batches_without_losing_runs
payload: Final = json.loads(request.prompt)
return ModelResult(content=json.dumps({"candidates": payload["candidates"]}), cost=0)
async def progress(_stage: str, _coverage: Coverage) -> None:
async def progress(_stage: str, _coverage: Coverage, _review: Review | None = None, /) -> None:
pass
result: Final = await cluster_batches(observation_batches(observations), model, progress, Coverage())
@ -950,7 +951,7 @@ async def test_invalid_candidate_response_preserves_other_findings_and_reports_i
return ModelResult(content="not JSON", cost=0)
return ModelResult(content=json.dumps({"action": "submit", "finding": finding("run").model_dump()}), cost=0)
async def progress(_stage: str, coverage: Coverage) -> None:
async def progress(_stage: str, coverage: Coverage, _review: Review | None = None, /) -> None:
counts.put(coverage.inconclusive)
claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=())
@ -1196,3 +1197,106 @@ async def test_investigator_can_read_all_evidence_pages_across_successive_span_b
assert result.error == ""
assert tuple(seen.get_nowait() for _ in range(seen.qsize())) == tuple(p.span_id for p in parts)
assert tuple(read_cursors.get_nowait() for _ in range(read_cursors.qsize())) == ("", "span039")
def test_review_flags_only_spans_cited_by_this_runs_observations() -> None:
from litellm.proxy.lens.analysis import Observation, review_of
execution: Final = Execution(
id="run1",
source="traces",
trace_id="trace-1",
team_id="",
name="task",
start_time="",
span_count=3,
service="bot",
)
shown: Final = tuple(
TracePart(execution_id="run1", span_id=span, name=span, kind="tool", content=f"{span} output")
for span in ("root", "search", "answer")
)
observation: Final = Observation(
check_id="retries",
summary="Search failed twice",
evidence=(
Evidence(execution_id="run1", span_id="search", quote="search output"),
Evidence(execution_id="other", span_id="answer", quote="answer output"),
),
)
examined: Final = Examined(
execution=execution,
observations=(observation,),
parts=shown,
partial=False,
cannot_assess=False,
reasoning="Asked to search; it retried without recovering.",
shown=shown,
)
review: Final = review_of(examined, "cerebras/model", 42, NOW)
assert tuple((s.span_id, s.cited) for s in review.spans) == (("root", False), ("search", True), ("answer", False))
assert (review.agent, review.trace_id, review.duration_ms) == ("bot", "trace-1", 42)
assert review.reasoning == examined.reasoning
assert tuple((v.check_id, v.summary) for v in review.verdicts) == (("retries", "Search failed twice"),)
@pytest.mark.asyncio
async def test_each_screened_run_reports_a_review_with_the_models_reasoning() -> None:
from litellm.proxy.lens.analysis import analyze_sample
execution: Final = Execution(
id="opaque-original", source="traces", trace_id="trace", team_id="", name="task", start_time="", span_count=2
)
reasoning: Final = "The user asked for a refund; the tool timed out and the agent gave up."
reviews: Final = SimpleQueue[Review]()
async def read(identity: str, _cursor: str, _offset: int) -> ExecutionContent:
return ExecutionContent(
execution=execution,
parts=(
TracePart(execution_id=identity, span_id="a-root", name="agent", kind="agent", content="Refund please"),
TracePart(
execution_id=identity,
span_id="b-tool",
parent_span_id="a-root",
name="refund",
kind="tool",
content="Tool timeout",
),
),
)
async def model(request: ModelRequest) -> ModelResult:
if request.purpose == "cluster":
return ModelResult(content='{"candidates":[]}', cost=0)
if request.purpose == "investigate":
return ModelResult(content='{"action":"inconclusive"}', cost=0)
return ModelResult(
content=json.dumps(
{
"reasoning": reasoning,
"observations": [
{
"check_id": "retries",
"summary": "Gave up after a timeout",
"evidence": [{"execution_id": "r0", "span_id": "b-tool", "quote": "Tool timeout"}],
}
],
}
),
cost=0,
)
async def progress(_stage: str, _coverage: Coverage, review: Review | None = None, /) -> None:
if review is not None:
reviews.put(review)
claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=())
await analyze_sample(claim, Sample(executions=(execution,), eligible=1), read, model, progress)
review: Final = reviews.get_nowait()
assert reviews.empty()
assert review.execution_id == execution.id
assert review.reasoning == reasoning
assert review.model == claim.job.settings.model
assert tuple((s.span_id, s.cited) for s in review.spans) == (("a-root", False), ("b-tool", True))
assert tuple(v.summary for v in review.verdicts) == ("Gave up after a timeout",)

View file

@ -12,6 +12,7 @@ from litellm.proxy.lens.models import (
ExecutionContent,
ModelRequest,
ModelResult,
Progress,
Result,
Sample,
TracePart,
@ -417,3 +418,43 @@ async def test_transient_heartbeat_failure_recovers_without_cancelling_analysis(
assert result.error == ""
assert result.coverage.screened == 1 and result.coverage.unassessable == 0
assert attempts.qsize() == 2 and saved.empty()
@pytest.mark.asyncio
async def test_worker_sends_each_runs_review_with_its_progress() -> 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
)
sent: Final = SimpleQueue[Progress]()
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="s", name="step", kind="agent", content="Done"),),
).model_dump(),
)
case "model":
return httpx.Response(
200, json={"content": '{"observations":[],"reasoning":"Finished the task."}', "cost": 0}
)
case "progress":
sent.put(Progress.model_validate_json(request.content))
return httpx.Response(200, json=True)
case "result":
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()
reviews: Final = tuple(p.review for p in (sent.get_nowait() for _ in range(sent.qsize())) if p.review)
assert tuple((r.execution_id, r.reasoning) for r in reviews) == (("run", "Finished the task."),)