mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
perf(lens): classify signals within seconds of a trace finishing (#45186)
* perf(lens): add a 2 second live signal sweep over recently finished traces Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> * perf(lens): run the live and backlog signal sweeps side by side Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> * test(lens): cover the live sweep window and backlog draining Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> * fix(lens): keep known signal results on screen and poll every 2s while runs wait Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> * test(lens): cover signal polling speed and results surviving list changes Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> * fix(lens): show Checking instead of Queued and leave clean runs blank Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> * test(lens): expect Checking for runs waiting on signals Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
parent
ec0af8d5f8
commit
9e20264607
7 changed files with 298 additions and 81 deletions
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import hashlib
|
||||
import json
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from itertools import accumulate
|
||||
from types import MappingProxyType
|
||||
|
|
@ -15,7 +16,7 @@ from litellm.litellm_core_utils.secret_redaction import redact_internal_details
|
|||
from litellm.proxy.lens.models import ActivitySelection, Execution, Record, Scope, TraceIdentity
|
||||
from litellm.proxy.lens.sources import SourceReader, Storage
|
||||
|
||||
SIGNAL_INTERVAL_SECONDS: Final = 60
|
||||
SIGNAL_SETTLE: Final = timedelta(seconds=15)
|
||||
SIGNAL_PAGE_SIZE: Final = 100
|
||||
SIGNAL_MAX_PER_TICK: Final = 50
|
||||
SIGNAL_CONCURRENCY: Final = 8
|
||||
|
|
@ -30,6 +31,19 @@ SIGNAL_TRANSCRIPT_MAX_CHARS: Final = 40000
|
|||
SIGNAL_TRANSCRIPT_HEAD_CHARS: Final = 15000
|
||||
SIGNAL_TRANSCRIPT_TAIL_CHARS: Final = 25000
|
||||
SIGNAL_MAX_SCAN_PAGES: Final = 10
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SignalSweep:
|
||||
lookback: timedelta
|
||||
interval_seconds: float
|
||||
max_pages: int
|
||||
|
||||
|
||||
SIGNAL_LIVE_SWEEP: Final = SignalSweep(lookback=timedelta(minutes=15), interval_seconds=2, max_pages=1)
|
||||
SIGNAL_BACKLOG_SWEEP: Final = SignalSweep(
|
||||
lookback=timedelta(hours=24), interval_seconds=60, max_pages=SIGNAL_MAX_SCAN_PAGES
|
||||
)
|
||||
SIGNAL_TASK: Final = (
|
||||
"An AI agent run recorded as a trace. Judge only what the user and the agent said and did in these steps."
|
||||
)
|
||||
|
|
@ -441,6 +455,7 @@ class _SignalScan:
|
|||
now: datetime,
|
||||
cursor: str,
|
||||
limit: int,
|
||||
sweep: SignalSweep,
|
||||
) -> None:
|
||||
self.reader: Final = reader
|
||||
self.repository: Final = repository
|
||||
|
|
@ -449,6 +464,7 @@ class _SignalScan:
|
|||
self.now: Final = now
|
||||
self.cursor: str = cursor
|
||||
self.limit: Final = limit
|
||||
self.sweep: Final = sweep
|
||||
self.executions: tuple[Execution, ...] = ()
|
||||
self.finished: bool = False
|
||||
|
||||
|
|
@ -478,9 +494,9 @@ class _SignalScan:
|
|||
return eligible, next_cursor
|
||||
|
||||
async def run(self) -> tuple[tuple[Execution, ...], str]:
|
||||
start: Final = int((self.now - timedelta(hours=24)).timestamp() * 1000)
|
||||
end: Final = int((self.now - timedelta(minutes=2)).timestamp() * 1000)
|
||||
for _ in range(SIGNAL_MAX_SCAN_PAGES):
|
||||
start: Final = int((self.now - self.sweep.lookback).timestamp() * 1000)
|
||||
end: Final = int((self.now - SIGNAL_SETTLE).timestamp() * 1000)
|
||||
for _ in range(self.sweep.max_pages):
|
||||
if self.finished or len(self.executions) >= self.limit:
|
||||
break
|
||||
eligible, next_cursor = await self._read_page(start, end)
|
||||
|
|
@ -501,13 +517,20 @@ async def _scan_pages(
|
|||
now: datetime,
|
||||
cursor: str,
|
||||
remaining: int,
|
||||
sweep: SignalSweep,
|
||||
) -> tuple[tuple[Execution, ...], str]:
|
||||
if remaining <= 0:
|
||||
return (), cursor
|
||||
scan: Final = _SignalScan(reader, repository, scope, config, now, cursor, remaining)
|
||||
scan: Final = _SignalScan(reader, repository, scope, config, now, cursor, remaining, sweep)
|
||||
return await scan.run()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SignalTick:
|
||||
cursor: str
|
||||
claimed: int = 0
|
||||
|
||||
|
||||
async def run_signal_tick(
|
||||
storage: Storage,
|
||||
repository: SignalRepositoryProtocol | None,
|
||||
|
|
@ -515,13 +538,14 @@ async def run_signal_tick(
|
|||
clock: Clock,
|
||||
router_ready: RouterReady = lambda: True,
|
||||
cursor: str = "",
|
||||
) -> str:
|
||||
sweep: SignalSweep = SIGNAL_BACKLOG_SWEEP,
|
||||
) -> SignalTick:
|
||||
if repository is None or completion is None or not router_ready():
|
||||
return cursor
|
||||
return SignalTick(cursor)
|
||||
now: Final = clock()
|
||||
config: Final = await repository.get_config()
|
||||
if not config.enabled:
|
||||
return cursor
|
||||
return SignalTick(cursor)
|
||||
reader: Final = SourceReader(storage)
|
||||
scope: Final = Scope(all_teams=True)
|
||||
candidates: Final = await _scan_pages(
|
||||
|
|
@ -532,12 +556,13 @@ async def run_signal_tick(
|
|||
now,
|
||||
cursor,
|
||||
SIGNAL_MAX_PER_TICK,
|
||||
sweep,
|
||||
)
|
||||
executions, next_cursor = candidates
|
||||
classifier: Final = SignalClassifier(reader, completion, clock)
|
||||
semaphore: Final = asyncio.Semaphore(SIGNAL_CONCURRENCY)
|
||||
|
||||
async def process(execution: Execution) -> None:
|
||||
async def process(execution: Execution) -> bool:
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
async with semaphore:
|
||||
|
|
@ -547,18 +572,37 @@ async def run_signal_tick(
|
|||
claimed: Final = await repository.claim(execution, config, claimed_until, claimed_at)
|
||||
except Exception as error:
|
||||
verbose_proxy_logger.error("Lens signal claim failed: %s", redact_internal_details(str(error)))
|
||||
return
|
||||
return False
|
||||
if not claimed:
|
||||
return
|
||||
return False
|
||||
await _process_claimed(classifier, repository, scope, execution, config, claimed_until)
|
||||
return True
|
||||
|
||||
await asyncio.gather(*(process(execution) for execution in executions))
|
||||
return next_cursor
|
||||
outcomes: Final = await asyncio.gather(*(process(execution) for execution in executions))
|
||||
return SignalTick(next_cursor, sum(outcomes))
|
||||
|
||||
|
||||
async def _logged_tick(
|
||||
storage: Storage,
|
||||
repository: SignalRepositoryProtocol | None,
|
||||
completion: DecisionsCall | None,
|
||||
clock: Clock,
|
||||
router_ready: RouterReady,
|
||||
cursor: str,
|
||||
sweep: SignalSweep,
|
||||
) -> SignalTick:
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
try:
|
||||
return await run_signal_tick(storage, repository, completion, clock, router_ready, cursor, sweep)
|
||||
except Exception as error:
|
||||
verbose_proxy_logger.error("Lens signal tick failed: %s", redact_internal_details(str(error)))
|
||||
return SignalTick(cursor)
|
||||
|
||||
|
||||
class _SignalLoopState:
|
||||
def __init__(self) -> None:
|
||||
self.cursor: str = ""
|
||||
self.tick: SignalTick = SignalTick("")
|
||||
|
||||
|
||||
async def run_signal_loop(
|
||||
|
|
@ -567,20 +611,9 @@ async def run_signal_loop(
|
|||
completion: DecisionsCall | None,
|
||||
clock: Clock = lambda: datetime.now(timezone.utc),
|
||||
router_ready: RouterReady = lambda: True,
|
||||
sweep: SignalSweep = SIGNAL_BACKLOG_SWEEP,
|
||||
) -> None:
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
state: Final = _SignalLoopState()
|
||||
while True:
|
||||
try:
|
||||
state.cursor = await run_signal_tick(
|
||||
storage,
|
||||
repository,
|
||||
completion,
|
||||
clock,
|
||||
router_ready,
|
||||
cursor=state.cursor,
|
||||
)
|
||||
except Exception as error:
|
||||
verbose_proxy_logger.error("Lens signal tick failed: %s", redact_internal_details(str(error)))
|
||||
await asyncio.sleep(SIGNAL_INTERVAL_SECONDS)
|
||||
state.tick = await _logged_tick(storage, repository, completion, clock, router_ready, state.tick.cursor, sweep)
|
||||
await asyncio.sleep(0 if state.tick.claimed >= SIGNAL_MAX_PER_TICK else sweep.interval_seconds)
|
||||
|
|
|
|||
|
|
@ -584,6 +584,8 @@ from litellm.proxy.lens.endpoints import router as lens_router
|
|||
from litellm.proxy.lens.repository import WriterDatabase
|
||||
from litellm.proxy.lens.signal_repository import SignalRepository
|
||||
from litellm.proxy.lens.signals import (
|
||||
SIGNAL_BACKLOG_SWEEP,
|
||||
SIGNAL_LIVE_SWEEP,
|
||||
DecisionQuestions,
|
||||
DecisionsCall,
|
||||
DecisionState,
|
||||
|
|
@ -1675,17 +1677,21 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState
|
|||
from litellm.proxy.admin_mcp import admin_mcp_lifespan
|
||||
|
||||
signal_completion: Final[DecisionsCall] = _call_current_lens_signal_router
|
||||
signal_task: Final = (
|
||||
asyncio.create_task(
|
||||
run_signal_loop(
|
||||
receiver.storage,
|
||||
SignalRepository(WriterDatabase(writer_wrapper(prisma_client.db))),
|
||||
signal_completion,
|
||||
router_ready=lambda: llm_router is not None,
|
||||
signal_tasks: Final = (
|
||||
tuple(
|
||||
asyncio.create_task(
|
||||
run_signal_loop(
|
||||
receiver.storage,
|
||||
SignalRepository(WriterDatabase(writer_wrapper(prisma_client.db))),
|
||||
signal_completion,
|
||||
router_ready=lambda: llm_router is not None,
|
||||
sweep=sweep,
|
||||
)
|
||||
)
|
||||
for sweep in (SIGNAL_LIVE_SWEEP, SIGNAL_BACKLOG_SWEEP)
|
||||
)
|
||||
if receiver is not None and prisma_client is not None
|
||||
else None
|
||||
else ()
|
||||
)
|
||||
|
||||
try:
|
||||
|
|
@ -1694,9 +1700,9 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState
|
|||
await admin_mcp_stack.enter_async_context(admin_mcp_lifespan(app))
|
||||
yield state
|
||||
finally:
|
||||
if signal_task is not None:
|
||||
for signal_task in signal_tasks:
|
||||
signal_task.cancel()
|
||||
await asyncio.gather(signal_task, return_exceptions=True)
|
||||
await asyncio.gather(*signal_tasks, return_exceptions=True)
|
||||
|
||||
if model_info_scheduler is not None and model_info_scheduler.running:
|
||||
model_info_scheduler.remove_job("refresh_model_info")
|
||||
|
|
|
|||
|
|
@ -15,7 +15,10 @@ from litellm.proxy.lens.repository import Database, Row
|
|||
from litellm.proxy.lens.signal_repository import SignalRepository
|
||||
from litellm.proxy.lens.signals import (
|
||||
DEFAULT_SIGNALS,
|
||||
SIGNAL_BACKLOG_SWEEP,
|
||||
SIGNAL_CLAIM_LEASE,
|
||||
SIGNAL_LIVE_SWEEP,
|
||||
SIGNAL_MAX_PER_TICK,
|
||||
SIGNAL_MAX_SCAN_PAGES,
|
||||
SIGNAL_TASK,
|
||||
DecisionQuestions,
|
||||
|
|
@ -26,6 +29,7 @@ from litellm.proxy.lens.signals import (
|
|||
SignalConfig,
|
||||
SignalData,
|
||||
SignalStep,
|
||||
SignalSweep,
|
||||
StoredTraceSignal,
|
||||
candidate,
|
||||
run_signal_loop,
|
||||
|
|
@ -711,15 +715,17 @@ async def test_signal_tick_resumes_after_ten_pages_and_resets_after_a_short_page
|
|||
) -> object:
|
||||
return {"answers": {}}
|
||||
|
||||
first_cursor: Final = await run_signal_tick(storage, repository, decide, lambda: NOW)
|
||||
first_cursor: Final = (await run_signal_tick(storage, repository, decide, lambda: NOW)).cursor
|
||||
first_calls: Final = tuple(storage.cursors.get_nowait() for _ in range(storage.cursors.qsize()))
|
||||
second_cursor: Final = await run_signal_tick(
|
||||
storage,
|
||||
repository,
|
||||
decide,
|
||||
lambda: NOW,
|
||||
cursor=first_cursor,
|
||||
)
|
||||
second_cursor: Final = (
|
||||
await run_signal_tick(
|
||||
storage,
|
||||
repository,
|
||||
decide,
|
||||
lambda: NOW,
|
||||
cursor=first_cursor,
|
||||
)
|
||||
).cursor
|
||||
second_calls: Final = tuple(storage.cursors.get_nowait() for _ in range(storage.cursors.qsize()))
|
||||
|
||||
assert len(first_calls) == SIGNAL_MAX_SCAN_PAGES
|
||||
|
|
@ -730,12 +736,14 @@ async def test_signal_tick_resumes_after_ten_pages_and_resets_after_a_short_page
|
|||
|
||||
short_storage: Final = PagedSampleStorage((pages[0][:50],))
|
||||
short_database: Final = SignalDatabase(config, stored_rows=stored_rows[:50])
|
||||
short_cursor: Final = await run_signal_tick(
|
||||
short_storage,
|
||||
SignalRepository(short_database),
|
||||
decide,
|
||||
lambda: NOW,
|
||||
)
|
||||
short_cursor: Final = (
|
||||
await run_signal_tick(
|
||||
short_storage,
|
||||
SignalRepository(short_database),
|
||||
decide,
|
||||
lambda: NOW,
|
||||
)
|
||||
).cursor
|
||||
assert short_cursor == ""
|
||||
|
||||
|
||||
|
|
@ -778,26 +786,30 @@ async def test_signal_tick_resumes_a_partially_consumed_page() -> None:
|
|||
}
|
||||
|
||||
first_database: Final = SignalDatabase(config, stored_rows=initial_rows)
|
||||
first_cursor: Final = await run_signal_tick(
|
||||
storage,
|
||||
SignalRepository(first_database),
|
||||
decide,
|
||||
lambda: NOW,
|
||||
cursor=resume_cursor,
|
||||
)
|
||||
first_cursor: Final = (
|
||||
await run_signal_tick(
|
||||
storage,
|
||||
SignalRepository(first_database),
|
||||
decide,
|
||||
lambda: NOW,
|
||||
cursor=resume_cursor,
|
||||
)
|
||||
).cursor
|
||||
first_claims: Final = tuple(first_database.claims.get_nowait() for _ in range(first_database.claims.qsize()))
|
||||
|
||||
classified_first_rows: Final = tuple(
|
||||
stored_trace(CURRENT_CONFIG_KEY, trace_id=trace_id) for trace_id in first_claims
|
||||
)
|
||||
second_database: Final = SignalDatabase(config, stored_rows=(*initial_rows, *classified_first_rows))
|
||||
second_cursor: Final = await run_signal_tick(
|
||||
storage,
|
||||
SignalRepository(second_database),
|
||||
decide,
|
||||
lambda: NOW,
|
||||
cursor=first_cursor,
|
||||
)
|
||||
second_cursor: Final = (
|
||||
await run_signal_tick(
|
||||
storage,
|
||||
SignalRepository(second_database),
|
||||
decide,
|
||||
lambda: NOW,
|
||||
cursor=first_cursor,
|
||||
)
|
||||
).cursor
|
||||
second_claims: Final = tuple(second_database.claims.get_nowait() for _ in range(second_database.claims.qsize()))
|
||||
sample_cursors: Final = tuple(storage.cursors.get_nowait() for _ in range(storage.cursors.qsize()))
|
||||
expected_eligible: Final = frozenset(f"trace-{index}" for index in range(20, 100))
|
||||
|
|
@ -1000,3 +1012,87 @@ async def test_proxy_signal_call_resolves_the_current_router(monkeypatch: pytest
|
|||
monkeypatch.setattr(proxy_server, "llm_router", None)
|
||||
with pytest.raises(RuntimeError, match="router is not initialized"):
|
||||
await call_current_router()
|
||||
|
||||
|
||||
class RecordingSampleStorage(PagedSampleStorage):
|
||||
def __init__(self, pages: tuple[tuple[ExecutionRow, ...], ...]) -> None:
|
||||
super().__init__(pages)
|
||||
self.windows: Final[asyncio.Queue[tuple[int, int]]] = asyncio.Queue()
|
||||
|
||||
async def lens_sample(self, parameters: LensSampleParams) -> Sequence[ExecutionRow]:
|
||||
await self.windows.put((parameters.start, parameters.end))
|
||||
return await super().lens_sample(parameters)
|
||||
|
||||
|
||||
def sample_rows(prefix: str, count: int) -> tuple[ExecutionRow, ...]:
|
||||
return tuple(
|
||||
ExecutionRow(
|
||||
source="traces",
|
||||
trace_id=f"{prefix}-{index}",
|
||||
team_id="",
|
||||
name=f"{prefix}-{index}",
|
||||
start_time="",
|
||||
span_count=1,
|
||||
root_seen=1,
|
||||
eligible=count,
|
||||
selected=count,
|
||||
selection_key=f"{prefix}-{index}",
|
||||
)
|
||||
for index in range(count)
|
||||
)
|
||||
|
||||
|
||||
async def no_answers(
|
||||
*,
|
||||
model: str,
|
||||
state: DecisionState,
|
||||
questions: DecisionQuestions,
|
||||
timeout: float,
|
||||
metadata: Mapping[str, object],
|
||||
) -> object:
|
||||
return {"answers": {}}
|
||||
|
||||
|
||||
def drained(queue: "asyncio.Queue[tuple[int, int]]") -> tuple[tuple[int, int], ...]:
|
||||
return tuple(queue.get_nowait() for _ in range(queue.qsize()))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_live_sweep_reads_one_page_of_recently_finished_traces() -> None:
|
||||
pages: Final = (sample_rows("a", 100), sample_rows("b", 100), ())
|
||||
stored_rows: Final = tuple(
|
||||
stored_trace(CURRENT_CONFIG_KEY, trace_id=row.trace_id) for row in chain.from_iterable(pages)
|
||||
)
|
||||
live_storage: Final = RecordingSampleStorage(pages)
|
||||
backlog_storage: Final = RecordingSampleStorage(pages)
|
||||
repository: Final = SignalRepository(SignalDatabase(SignalConfig(model="decision"), stored_rows=stored_rows))
|
||||
|
||||
live_tick: Final = await run_signal_tick(live_storage, repository, no_answers, lambda: NOW, sweep=SIGNAL_LIVE_SWEEP)
|
||||
await run_signal_tick(backlog_storage, repository, no_answers, lambda: NOW, sweep=SIGNAL_BACKLOG_SWEEP)
|
||||
live_windows: Final = drained(live_storage.windows)
|
||||
backlog_windows: Final = drained(backlog_storage.windows)
|
||||
now_ms: Final = int(NOW.timestamp() * 1000)
|
||||
|
||||
assert len(live_windows) == 1
|
||||
assert live_tick.cursor == pages[0][-1].selection_key
|
||||
assert live_tick.claimed == 0
|
||||
assert backlog_windows[0][0] < live_windows[0][0] < live_windows[0][1] < now_ms
|
||||
assert live_windows[0][1] == backlog_windows[0][1]
|
||||
assert now_ms - live_windows[0][1] <= 30_000, "a finished trace should be visible to the sweep within seconds"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_signal_loop_drains_a_backlog_without_waiting_for_the_interval() -> None:
|
||||
storage: Final = SignalStorage(executions=sample_rows("trace", SIGNAL_MAX_PER_TICK + 10))
|
||||
database: Final = SignalDatabase(SignalConfig(model="decision"))
|
||||
hour_long_sweep: Final = SignalSweep(lookback=timedelta(minutes=15), interval_seconds=3600, max_pages=1)
|
||||
|
||||
task: Final = asyncio.create_task(
|
||||
run_signal_loop(storage, SignalRepository(database), no_answers, lambda: NOW, sweep=hour_long_sweep)
|
||||
)
|
||||
claims: Final = tuple([await asyncio.wait_for(database.claims.get(), 1) for _ in range(SIGNAL_MAX_PER_TICK + 1)])
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
assert len(claims) == SIGNAL_MAX_PER_TICK + 1
|
||||
|
|
|
|||
|
|
@ -317,7 +317,7 @@ describe("AgentTracesTable signals", () => {
|
|||
).toEqual(["User frustration", "Repeated request"]);
|
||||
expect(flagged).toHaveAttribute("title", "Signals: User frustration (92%), Repeated request (71%)");
|
||||
expect(within(rows[1]).getByTitle("No signals detected")).toBeInTheDocument();
|
||||
expect(within(rows[2]).getByText("Queued")).toBeInTheDocument();
|
||||
expect(within(rows[2]).getByText("Checking")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("keeps the signals column with a setup link until signals are configured", async () => {
|
||||
|
|
|
|||
|
|
@ -170,11 +170,11 @@ function SignalsCell({ run }: { run: TraceSummary }) {
|
|||
if (!state || state.status === "pending") return <Skeleton aria-label="Loading signals" className="h-3 w-16" />;
|
||||
if (state.status === "error") return muted("Unavailable", "Could not load signals");
|
||||
const { status } = state.signals;
|
||||
if (status === "unclassified") return muted("Queued", "Waiting for the System 1 model to check this run");
|
||||
if (status === "pending") return muted("Checking", "The System 1 model is checking this run");
|
||||
if (status === "unclassified" || status === "pending")
|
||||
return muted("Checking", "The System 1 model is checking this run");
|
||||
if (status === "failed") return muted("Not checked", "The System 1 model could not check this run");
|
||||
const flags = flaggedSignals(state.signals);
|
||||
if (!flags.length) return muted("-", "No signals detected");
|
||||
if (!flags.length) return <span title="No signals detected" />;
|
||||
return <SignalPills flags={flags} className="overflow-hidden" />;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,59 @@
|
|||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
import { renderHook, waitFor } from "@testing-library/react";
|
||||
import type { PropsWithChildren } from "react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
|
||||
import { TracesApiContext, type TracesApi } from "../api";
|
||||
import type { TraceSignals, TraceSummary } from "../types";
|
||||
import { signalPollInterval, useTraceSignals } from "./useTraceSignals";
|
||||
|
||||
const run = (trace_id: string): TraceSummary => ({ trace_id, trace_ref: "" }) as TraceSummary;
|
||||
|
||||
const signals = (trace_id: string, status: TraceSignals["status"]): TraceSignals => ({
|
||||
trace_id,
|
||||
trace_ref: "",
|
||||
status,
|
||||
flags: status === "classified" ? [{ signal_id: "tool_failure", name: "Tool failure", score: 0.9 }] : [],
|
||||
model: "jev",
|
||||
classified_at: null,
|
||||
});
|
||||
|
||||
describe("signalPollInterval", () => {
|
||||
it("polls fast only while a run is still waiting for a result", () => {
|
||||
const settled = signalPollInterval([signals("a", "classified"), signals("b", "failed")]);
|
||||
expect(signalPollInterval([signals("a", "classified"), signals("b", "unclassified")])).toBeLessThan(settled);
|
||||
expect(signalPollInterval([signals("a", "pending")])).toBeLessThan(settled);
|
||||
expect(signalPollInterval(undefined)).toBe(settled);
|
||||
});
|
||||
});
|
||||
|
||||
describe("useTraceSignals", () => {
|
||||
it("keeps showing known results while a new run is added to the list", async () => {
|
||||
const gate = { release: (): void => undefined };
|
||||
const api = {
|
||||
live: true,
|
||||
signals: vi.fn(async (traces: { trace_id: string }[]) => {
|
||||
if (traces.length > 1) await new Promise<void>((resolve) => (gate.release = resolve));
|
||||
return traces.map((trace) => signals(trace.trace_id, "classified"));
|
||||
}),
|
||||
} as unknown as TracesApi;
|
||||
const client = new QueryClient({ defaultOptions: { queries: { retry: false } } });
|
||||
const wrapper = ({ children }: PropsWithChildren) => (
|
||||
<QueryClientProvider client={client}>
|
||||
<TracesApiContext.Provider value={api}>{children}</TracesApiContext.Provider>
|
||||
</QueryClientProvider>
|
||||
);
|
||||
const { result, rerender } = renderHook(({ runs }) => useTraceSignals("token", runs, true), {
|
||||
wrapper,
|
||||
initialProps: { runs: [run("old")] },
|
||||
});
|
||||
await waitFor(() => expect(result.current.get("old")?.status).toBe("ready"));
|
||||
|
||||
rerender({ runs: [run("new"), run("old")] });
|
||||
|
||||
expect(result.current.get("old")?.status).toBe("ready");
|
||||
expect(result.current.get("new")?.status).toBe("pending");
|
||||
gate.release();
|
||||
await waitFor(() => expect(result.current.get("new")?.status).toBe("ready"));
|
||||
});
|
||||
});
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
import { useQueries, useQuery } from "@tanstack/react-query";
|
||||
import { useQueries, useQuery, useQueryClient, type Query, type QueryClient } from "@tanstack/react-query";
|
||||
import { chunk } from "es-toolkit";
|
||||
|
||||
import { useTracesApi } from "../api";
|
||||
|
|
@ -6,13 +6,34 @@ import type { SignalFlag, TraceSignals, TraceSummary } from "../types";
|
|||
|
||||
export type TraceSignalState = { status: "ready"; signals: TraceSignals } | { status: "pending" } | { status: "error" };
|
||||
|
||||
const POLL_MS = 15000;
|
||||
const IDLE_POLL_MS = 15000;
|
||||
const ACTIVE_POLL_MS = 2000;
|
||||
|
||||
const identity = ({ trace_id, trace_ref }: { trace_id: string; trace_ref?: string | null }) => ({
|
||||
trace_id,
|
||||
trace_ref: trace_ref ?? "",
|
||||
});
|
||||
|
||||
const awaitingResult = (signals: TraceSignals): boolean =>
|
||||
signals.status === "unclassified" || signals.status === "pending";
|
||||
|
||||
export const signalPollInterval = (results: readonly TraceSignals[] | undefined): number =>
|
||||
results?.some(awaitingResult) ? ACTIVE_POLL_MS : IDLE_POLL_MS;
|
||||
|
||||
const pollWhile = (enabled: boolean, live: boolean) => (query: Query<TraceSignals[]>) =>
|
||||
enabled && live ? signalPollInterval(query.state.data) : false;
|
||||
|
||||
const signalKey = (result: { trace_id: string; trace_ref?: string | null }): string =>
|
||||
result.trace_ref || result.trace_id;
|
||||
|
||||
const cachedSignals = (client: QueryClient, accessToken: string): Map<string, TraceSignals> =>
|
||||
new Map(
|
||||
client
|
||||
.getQueriesData<TraceSignals[]>({ queryKey: ["traceSignals", accessToken] })
|
||||
.flatMap(([, data]) => data ?? [])
|
||||
.map((result) => [signalKey(result), result]),
|
||||
);
|
||||
|
||||
export const flaggedSignals = (signals?: TraceSignals): SignalFlag[] =>
|
||||
signals?.status === "classified" ? signals.flags ?? [] : [];
|
||||
|
||||
|
|
@ -21,28 +42,30 @@ export const isFlagged = (state?: TraceSignalState): boolean =>
|
|||
|
||||
export function useTraceSignals(accessToken: string, runs: TraceSummary[], enabled: boolean) {
|
||||
const api = useTracesApi(accessToken);
|
||||
const client = useQueryClient();
|
||||
const batches = chunk(runs.map(identity), 500);
|
||||
const queries = useQueries({
|
||||
queries: batches.map((traces) => ({
|
||||
queryKey: ["traceSignals", accessToken, traces],
|
||||
queryFn: () => api.signals(traces),
|
||||
enabled,
|
||||
staleTime: POLL_MS,
|
||||
refetchInterval: enabled && api.live ? POLL_MS : false,
|
||||
staleTime: ACTIVE_POLL_MS,
|
||||
refetchInterval: pollWhile(enabled, api.live),
|
||||
retry: false,
|
||||
})),
|
||||
});
|
||||
const known = cachedSignals(client, accessToken);
|
||||
return new Map<string, TraceSignalState>(
|
||||
batches.flatMap((traces, index) => {
|
||||
const query = queries[index];
|
||||
const results = new Map(query.data?.map((result) => [result.trace_ref || result.trace_id, result]));
|
||||
const results = new Map(query.data?.map((result) => [signalKey(result), result]));
|
||||
return traces.map((trace): [string, TraceSignalState] => {
|
||||
const key = trace.trace_ref || trace.trace_id;
|
||||
const found = results.get(key);
|
||||
const key = signalKey(trace);
|
||||
const found = results.get(key) ?? known.get(key);
|
||||
if (found) return [key, { status: "ready", signals: found }];
|
||||
if (query.isError) return [key, { status: "error" }];
|
||||
if (query.isPending) return [key, { status: "pending" }];
|
||||
if (!found) return [key, { status: "error" }];
|
||||
return [key, { status: "ready", signals: found }];
|
||||
return [key, { status: "error" }];
|
||||
});
|
||||
}),
|
||||
);
|
||||
|
|
@ -59,8 +82,8 @@ export function useTraceSignalFlags(
|
|||
queryKey: ["traceSignals", accessToken, traces],
|
||||
queryFn: () => api.signals(traces),
|
||||
enabled,
|
||||
staleTime: POLL_MS,
|
||||
refetchInterval: enabled && api.live ? POLL_MS : (false as const),
|
||||
staleTime: ACTIVE_POLL_MS,
|
||||
refetchInterval: pollWhile(enabled, api.live),
|
||||
retry: false,
|
||||
};
|
||||
const query = useQuery<TraceSignals[]>(queryOptions);
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue