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);