diff --git a/litellm/proxy/lens/signals.py b/litellm/proxy/lens/signals.py
index cfe273eb6af..9a25ab3eed5 100644
--- a/litellm/proxy/lens/signals.py
+++ b/litellm/proxy/lens/signals.py
@@ -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)
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index 68e10467c63..de420ba39e7 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -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")
diff --git a/tests/unit/proxy/lens/test_signals.py b/tests/unit/proxy/lens/test_signals.py
index d2c4eda3966..c70fc4b4f28 100644
--- a/tests/unit/proxy/lens/test_signals.py
+++ b/tests/unit/proxy/lens/test_signals.py
@@ -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
diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx
index 1bcb0ed1048..b2a8111656f 100644
--- a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx
+++ b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx
@@ -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 () => {
diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.tsx
index b99edde382d..fa4051a34a4 100644
--- a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.tsx
+++ b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.tsx
@@ -170,11 +170,11 @@ function SignalsCell({ run }: { run: TraceSummary }) {
if (!state || state.status === "pending") return ;
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 ;
return ;
}
diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/useTraceSignals.test.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/useTraceSignals.test.tsx
new file mode 100644
index 00000000000..c85bf144cef
--- /dev/null
+++ b/ui/litellm-dashboard/src/components/lens/traces/list/useTraceSignals.test.tsx
@@ -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((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) => (
+
+ {children}
+
+ );
+ 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"));
+ });
+});
diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/useTraceSignals.ts b/ui/litellm-dashboard/src/components/lens/traces/list/useTraceSignals.ts
index b798d0e4024..024519e4444 100644
--- a/ui/litellm-dashboard/src/components/lens/traces/list/useTraceSignals.ts
+++ b/ui/litellm-dashboard/src/components/lens/traces/list/useTraceSignals.ts
@@ -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) =>
+ 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 =>
+ new Map(
+ client
+ .getQueriesData({ 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(
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(queryOptions);