diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000200_lens_trace_signals/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000200_lens_trace_signals/migration.sql new file mode 100644 index 00000000000..67043eccb5f --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000200_lens_trace_signals/migration.sql @@ -0,0 +1,16 @@ +CREATE TABLE IF NOT EXISTS "LiteLLM_LensSignalConfig" ( + "id" TEXT NOT NULL, + "data" JSONB NOT NULL, + CONSTRAINT "LiteLLM_LensSignalConfig_pkey" PRIMARY KEY ("id") +); + +CREATE TABLE IF NOT EXISTS "LiteLLM_LensTraceSignal" ( + "trace_id" TEXT NOT NULL, + "trace_ref" TEXT NOT NULL DEFAULT '', + "config_key" TEXT NOT NULL, + "span_count" INTEGER NOT NULL, + "claimed_until" TIMESTAMP(3), + "classified_at" TIMESTAMP(3), + "data" JSONB NOT NULL, + CONSTRAINT "LiteLLM_LensTraceSignal_pkey" PRIMARY KEY ("trace_id", "trace_ref") +); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index ceb127c31bd..3b83c5b09cc 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1977,3 +1977,20 @@ model LiteLLM_LensDataset { @@id([id, revision]) } + +model LiteLLM_LensSignalConfig { + id String @id + data Json +} + +model LiteLLM_LensTraceSignal { + trace_id String + trace_ref String @default("") + config_key String + span_count Int + claimed_until DateTime? + classified_at DateTime? + data Json + + @@id([trace_id, trace_ref]) +} diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index 3b17aea3cd8..b22a8b9415b 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -353,6 +353,8 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset( "LiteLLM_LensRun", "LiteLLM_LensReview", "LiteLLM_LensWorker", + "LiteLLM_LensSignalConfig", + "LiteLLM_LensTraceSignal", ) ) PRISMA_RELATIONS: Final[frozenset[str]] = _PRISMA_MODELS | _PRISMA_VIEWS diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index d8b89a4a90e..a6381e10e39 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -52,6 +52,8 @@ from litellm.proxy.lens.models import ( from litellm.proxy.lens.release import PROTOCOL_VERSION, release_tag, worker_image from litellm.proxy.lens.repository import DueLens, LensRepository, WriterDatabase from litellm.proxy.lens.reviews import criteria_key +from litellm.proxy.lens.signal_repository import SignalRepository +from litellm.proxy.lens.signals import SignalConfig, TraceSignals, trace_signals from litellm.proxy.lens.sources import ActivityAvailability, SourceReader, Storage, parse_execution from litellm.proxy.lens.state import ( can_access, @@ -69,6 +71,7 @@ from litellm.proxy.lens.state import ( summarized, ) from litellm.proxy.tracing_runtime import provide_storage +from litellm.router import Router from litellm.types.llms.base import LiteLLMBaseModel router: Final = APIRouter(prefix="/lens", tags=["Lens"]) @@ -101,6 +104,14 @@ def repository() -> LensRepository: return LensRepository(WriterDatabase(writer_wrapper(prisma_client.db))) +def signals_repository() -> SignalRepository: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(503, "Lens needs a connected Postgres database") + return SignalRepository(WriterDatabase(writer_wrapper(prisma_client.db))) + + def source_reader(storage: Storage | None) -> SourceReader: if storage is None: raise HTTPException( @@ -118,6 +129,20 @@ def user_scope(auth: UserAPIKeyAuth, write: bool = False) -> Scope: raise HTTPException(403, "Lens requires proxy administrator access") +def validate_signal_model(config: SignalConfig, llm_router: Router | None) -> None: + if not config.model: + return + message: Final = "Choose a System 1 model (evaluation mode) configured on this proxy" + if llm_router is None: + raise HTTPException(400, message) + try: + model_group: Final = llm_router.get_model_group_info(model_group=config.model) + except Exception as error: + raise HTTPException(400, message) from error + if model_group is None or model_group.mode != "evaluation": + raise HTTPException(400, message) + + async def get_lens(lens_id: str, scope: Scope) -> Lens: lens: Final = await repository().get(lens_id) if lens is None or not can_access(scope, lens.scope): @@ -261,6 +286,39 @@ async def list_agents(auth: Auth, storage: StorageDep) -> tuple[str, ...]: return await source_reader(storage).agents(scope) if storage is not None else () +@router.get("/signals", response_model=SignalConfig) +async def get_signals(auth: Auth) -> SignalConfig: + user_scope(auth) + return await signals_repository().get_config() + + +@router.put("/signals", response_model=SignalConfig) +async def put_signals(body: SignalConfig, auth: Auth) -> SignalConfig: + user_scope(auth, write=True) + from litellm.proxy.proxy_server import llm_router + + validate_signal_model(body, llm_router) + await signals_repository().save_config(body) + return body + + +@router.post("/traces/signals", response_model=tuple[TraceSignals, ...]) +async def trace_signal_statuses(body: TraceFindingsRequest, auth: Auth) -> tuple[TraceSignals, ...]: + user_scope(auth) + repo: Final = signals_repository() + config: Final = await repo.get_config() + existing: Final = await repo.traces(body.traces) + rows: Final = MappingProxyType({(row.trace_id, row.trace_ref): row for row in existing}) + return tuple( + trace_signals( + trace, + rows.get((trace.trace_id, trace.trace_ref)), + config, + ) + for trace in body.traces + ) + + @router.post("/traces/findings", response_model=tuple[TraceFindingCount, ...]) async def trace_findings(body: TraceFindingsRequest, auth: Auth) -> tuple[TraceFindingCount, ...]: user_scope(auth) diff --git a/litellm/proxy/lens/signal_repository.py b/litellm/proxy/lens/signal_repository.py new file mode 100644 index 00000000000..462cbdf13f4 --- /dev/null +++ b/litellm/proxy/lens/signal_repository.py @@ -0,0 +1,142 @@ +import json +from datetime import datetime +from typing import Final + +from pydantic import TypeAdapter + +from litellm.proxy.lens.models import Execution, TraceIdentity +from litellm.proxy.lens.repository import Database, Row +from litellm.proxy.lens.signals import ( + SIGNAL_RECLASSIFY_AFTER, + SIGNAL_RETRY_FAILED_AFTER, + SignalAttempt, + SignalConfig, + StoredTraceSignal, +) + +_ROWS: Final[TypeAdapter[tuple[Row, ...]]] = TypeAdapter(tuple[Row, ...]) + + +class SignalRepository: + def __init__(self, db: Database) -> None: + self.db: Final = db + + async def get_config(self) -> SignalConfig: + rows: Final = _ROWS.validate_python( + await self.db.query_raw('SELECT data FROM "LiteLLM_LensSignalConfig" WHERE id=$1', "global") + ) + return SignalConfig() if not rows else SignalConfig.model_validate(rows[0].data) + + async def save_config(self, config: SignalConfig) -> None: + await self.db.execute_raw( + """INSERT INTO "LiteLLM_LensSignalConfig" (id, data) + VALUES ($1, $2::jsonb) + ON CONFLICT (id) DO UPDATE SET data=EXCLUDED.data""", + "global", + json.dumps(config.model_dump(mode="json")), + ) + + async def traces(self, identities: tuple[TraceIdentity, ...]) -> tuple[StoredTraceSignal, ...]: + if not identities: + return () + payload: Final = json.dumps( + tuple({"trace_id": trace.trace_id, "trace_ref": trace.trace_ref} for trace in identities) + ) + rows: Final = _ROWS.validate_python( + await self.db.query_raw( + """SELECT jsonb_build_object( + 'trace_id', trace_id, + 'trace_ref', trace_ref, + 'config_key', config_key, + 'span_count', span_count, + 'claimed_until', claimed_until, + 'classified_at', classified_at, + 'data', data + ) AS data + FROM "LiteLLM_LensTraceSignal" + WHERE (trace_id, trace_ref) IN ( + SELECT trace_id, trace_ref FROM jsonb_to_recordset($1::jsonb) AS requested( + trace_id text, trace_ref text + ) + )""", + payload, + ) + ) + return tuple(StoredTraceSignal.model_validate(row.data) for row in rows) + + async def claim( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + now: datetime, + ) -> bool: + data: Final = json.dumps({"status": "pending", "scores": {}, "model": config.model, "error": ""}) + rows: Final = _ROWS.validate_python( + await self.db.query_raw( + """INSERT INTO "LiteLLM_LensTraceSignal" AS stored + (trace_id, trace_ref, config_key, span_count, claimed_until, classified_at, data) + VALUES ($1, $2, $3, $4, $5::timestamp, NULL, $6::jsonb) + ON CONFLICT (trace_id, trace_ref) DO UPDATE SET + config_key=EXCLUDED.config_key, + span_count=EXCLUDED.span_count, + claimed_until=EXCLUDED.claimed_until, + classified_at=NULL, + data=EXCLUDED.data + WHERE (stored.claimed_until IS NULL OR stored.claimed_until < $7::timestamp) + AND ( + stored.config_key IS DISTINCT FROM EXCLUDED.config_key + OR ( + stored.data->>'status'='pending' + AND stored.claimed_until < $7::timestamp + ) + OR ( + EXCLUDED.span_count > stored.span_count + AND stored.classified_at < $8::timestamp + ) + OR ( + stored.data->>'status'='failed' + AND stored.classified_at < $9::timestamp + ) + ) + RETURNING jsonb_build_object('trace_id', trace_id) AS data""", + execution.trace_id, + execution.trace_ref, + config.key(), + execution.span_count, + claimed_until, + data, + now, + now - SIGNAL_RECLASSIFY_AFTER, + now - SIGNAL_RETRY_FAILED_AFTER, + ) + ) + return bool(rows) + + async def store( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + classified_at: datetime, + attempt: SignalAttempt, + ) -> None: + payload: Final = json.dumps( + { + "status": attempt.status, + "scores": dict(attempt.scores), + "model": attempt.model, + "error": attempt.error, + } + ) + await self.db.execute_raw( + """UPDATE "LiteLLM_LensTraceSignal" + SET classified_at=$1::timestamp, claimed_until=NULL, data=$2::jsonb + WHERE trace_id=$3 AND trace_ref=$4 AND config_key=$5 AND claimed_until=$6::timestamp""", + classified_at, + payload, + execution.trace_id, + execution.trace_ref, + config.key(), + claimed_until, + ) diff --git a/litellm/proxy/lens/signals.py b/litellm/proxy/lens/signals.py new file mode 100644 index 00000000000..cfe273eb6af --- /dev/null +++ b/litellm/proxy/lens/signals.py @@ -0,0 +1,586 @@ +import asyncio +import hashlib +import json +from collections.abc import Callable, Mapping +from datetime import datetime, timedelta, timezone +from itertools import accumulate +from types import MappingProxyType +from typing import Annotated, Final, Literal, Protocol, TypeAlias + +from pydantic import ConfigDict, Field, JsonValue, ValidationError, field_validator, model_validator + +from litellm.integrations.clickhouse.context import lens_analysis +from litellm.litellm_core_utils.initialize_dynamic_callback_params import inherit_message_logging_privacy +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_PAGE_SIZE: Final = 100 +SIGNAL_MAX_PER_TICK: Final = 50 +SIGNAL_CONCURRENCY: Final = 8 +SIGNAL_CLAIM_LEASE: Final = timedelta(minutes=5) +SIGNAL_RECLASSIFY_AFTER: Final = timedelta(minutes=5) +SIGNAL_RETRY_FAILED_AFTER: Final = timedelta(minutes=30) +SIGNAL_MAX_CONTENT_PAGES: Final = 3 +SIGNAL_PART_MAX_CHARS: Final = 2000 +SIGNAL_PART_HEAD_CHARS: Final = 800 +SIGNAL_PART_TAIL_CHARS: Final = 1200 +SIGNAL_TRANSCRIPT_MAX_CHARS: Final = 40000 +SIGNAL_TRANSCRIPT_HEAD_CHARS: Final = 15000 +SIGNAL_TRANSCRIPT_TAIL_CHARS: Final = 25000 +SIGNAL_MAX_SCAN_PAGES: Final = 10 +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." +) + + +class Signal(Record): + id: str = Field(pattern=r"^[a-z][a-z0-9_]{0,63}$") + name: str = Field(min_length=1, max_length=60) + question: str = Field(min_length=3, max_length=500) + + +DEFAULT_SIGNALS: Final[tuple[Signal, ...]] = ( + Signal( + id="user_frustration", + name="User frustration", + question=( + "Does the user show frustration, annoyance or dissatisfaction with the agent in this run, for example " + "complaints, irritated corrections, all caps, profanity, or giving up on the task?" + ), + ), + Signal( + id="missing_capability", + name="Missing capability", + question=( + "Does the user ask for something the agent cannot do in this run, so that the agent refuses, says it " + "lacks a tool, permission, integration or data source, or fails because the capability does not exist?" + ), + ), + Signal( + id="repeated_request", + name="Repeated request", + question=( + "Does the user ask for the same thing more than once in this run, usually because the agent did not " + "deliver it the first time?" + ), + ), +) + + +class SignalConfig(Record): + model: str = "" + threshold: float = Field(default=0.5, ge=0.05, le=0.95, allow_inf_nan=False) + signals: tuple[Signal, ...] = DEFAULT_SIGNALS + + @model_validator(mode="after") + def validate_signals(self) -> "SignalConfig": + if len(self.signals) > 20: + raise ValueError("A maximum of 20 signals is allowed") + if len(frozenset(signal.id for signal in self.signals)) != len(self.signals): + raise ValueError("Signal IDs must be unique") + return self + + @property + def enabled(self) -> bool: + return bool(self.model) and bool(self.signals) + + def key(self) -> str: + payload: Final = json.dumps( + { + "model": self.model, + "signals": tuple({"id": signal.id, "question": signal.question} for signal in self.signals), + }, + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ) + return hashlib.sha256(payload.encode()).hexdigest() + + +Score: TypeAlias = Annotated[float, Field(ge=0, le=1, allow_inf_nan=False)] + + +class SignalFlag(Record): + signal_id: str + name: str + score: Score + + +class TraceSignals(TraceIdentity): + status: Literal["unclassified", "pending", "classified", "failed"] + flags: tuple[SignalFlag, ...] = () + model: str = "" + classified_at: datetime | None = None + + +class SignalStep(Record): + kind: str + name: str + content: str + + +class SignalData(Record): + model_config = ConfigDict(extra="ignore") + + status: Literal["pending", "classified", "failed"] = "pending" + scores: Mapping[str, Score] = Field(default_factory=lambda: MappingProxyType({})) + model: str = "" + error: str = "" + + +class SignalAttempt(Record): + status: Literal["classified", "failed"] + scores: Mapping[str, Score] = Field(default_factory=lambda: MappingProxyType({})) + model: str + error: str = "" + + +class StoredTraceSignal(Record): + trace_id: str + trace_ref: str = "" + config_key: str + span_count: int + claimed_until: datetime | None = None + classified_at: datetime | None = None + data: JsonValue + + @field_validator("claimed_until", "classified_at") + @classmethod + def normalize_database_timestamp(cls, value: datetime | None) -> datetime | None: + if value is not None and value.tzinfo is None: + return value.replace(tzinfo=timezone.utc) + return value + + +class NoulAnswer(Record): + model_config = ConfigDict(extra="ignore", allow_inf_nan=False, from_attributes=True) + + type: Literal["noul"] + noul: float = Field(ge=0, le=1, allow_inf_nan=False) + + +class DecisionsOutput(Record): + model_config = ConfigDict(extra="ignore", from_attributes=True) + + answers: Mapping[str, object] + + +DecisionState: TypeAlias = Mapping[str, object] +DecisionQuestions: TypeAlias = Mapping[str, Mapping[str, str]] +Clock: TypeAlias = Callable[[], datetime] +RouterReady: TypeAlias = Callable[[], bool] + + +class DecisionsCall(Protocol): + async def __call__( + self, + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: ... + + +class SignalRepositoryProtocol(Protocol): + async def get_config(self) -> SignalConfig: ... + async def traces(self, identities: tuple[TraceIdentity, ...]) -> tuple[StoredTraceSignal, ...]: ... + async def claim( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + now: datetime, + ) -> bool: ... + async def store( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + classified_at: datetime, + attempt: SignalAttempt, + ) -> None: ... + + +def signal_identity(trace: TraceIdentity | StoredTraceSignal | Execution) -> tuple[str, str]: + return trace.trace_id, trace.trace_ref + + +def candidate( + trace: Execution, + existing: StoredTraceSignal | None, + config_key: str, + now: datetime, +) -> bool: + if existing is None: + return True + if existing.claimed_until is not None and existing.claimed_until > now: + return False + if existing.config_key != config_key: + return True + status: Final = existing.data.get("status") if isinstance(existing.data, dict) else "" + if status == "pending": + return existing.claimed_until is not None and existing.claimed_until <= now + if existing.span_count > trace.span_count: + return False + if existing.span_count < trace.span_count: + return existing.classified_at is not None and existing.classified_at < now - SIGNAL_RECLASSIFY_AFTER + return ( + status == "failed" + and existing.classified_at is not None + and existing.classified_at < now - SIGNAL_RETRY_FAILED_AFTER + ) + + +def _take_head(steps: tuple[SignalStep, ...], remaining: int) -> tuple[SignalStep, ...]: + if remaining <= 0: + return () + cumulative_lengths: Final = tuple(accumulate(len(step.content) for step in steps)) + boundary: Final = next((index for index, total in enumerate(cumulative_lengths) if total >= remaining), None) + if boundary is None: + return steps + preceding: Final = steps[:boundary] + last: Final = steps[boundary] + used: Final = cumulative_lengths[boundary - 1] if boundary > 0 else 0 + last_length: Final = remaining - used + return ( + *preceding, + last + if last_length == len(last.content) + else last.model_copy(update=MappingProxyType({"content": last.content[:last_length]})), + ) + + +def _take_tail(steps: tuple[SignalStep, ...], remaining: int) -> tuple[SignalStep, ...]: + if remaining <= 0: + return () + reversed_steps: Final = tuple(reversed(steps)) + cumulative_lengths: Final = tuple(accumulate(len(step.content) for step in reversed_steps)) + boundary: Final = next((index for index, total in enumerate(cumulative_lengths) if total >= remaining), None) + if boundary is None: + return steps + preceding: Final = reversed_steps[:boundary] + last: Final = reversed_steps[boundary] + used: Final = cumulative_lengths[boundary - 1] if boundary > 0 else 0 + last_length: Final = remaining - used + selected: Final = ( + *preceding, + last + if last_length == len(last.content) + else last.model_copy(update=MappingProxyType({"content": last.content[-last_length:]})), + ) + return tuple(reversed(selected)) + + +def _bounded_steps(steps: tuple[SignalStep, ...]) -> tuple[SignalStep, ...]: + if sum(len(step.content) for step in steps) <= SIGNAL_TRANSCRIPT_MAX_CHARS: + return steps + head: Final = _take_head(steps, SIGNAL_TRANSCRIPT_HEAD_CHARS) + tail: Final = _take_tail(steps, SIGNAL_TRANSCRIPT_TAIL_CHARS) + omitted_count: Final = len(steps) - len(head) - len(tail) + marker: Final = SignalStep(kind="omitted", name="", content=f"{omitted_count} steps omitted") + return (*head, marker, *tail) + + +def _part_excerpt(content: str) -> str: + if len(content) <= SIGNAL_PART_MAX_CHARS: + return content + omitted: Final = len(content) - SIGNAL_PART_MAX_CHARS + marker: Final = f"\n[... {omitted} characters omitted ...]\n" + return f"{content[:SIGNAL_PART_HEAD_CHARS]}{marker}{content[-SIGNAL_PART_TAIL_CHARS:]}" + + +async def _content_pages( + reader: SourceReader, + scope: Scope, + execution: Execution, + cursor: str, + pages_left: int, +) -> tuple[SignalStep, ...]: + if pages_left == 0: + return () + content: Final = await reader.content(scope, execution, cursor) + current: Final = tuple( + SignalStep(kind=part.kind, name=part.name, content=_part_excerpt(part.content)) for part in content.parts + ) + rest: Final = ( + await _content_pages(reader, scope, execution, content.next_cursor, pages_left - 1) + if content.next_cursor is not None + else () + ) + return (*current, *rest) + + +async def signal_state(reader: SourceReader, scope: Scope, execution: Execution) -> DecisionState: + steps: Final = _bounded_steps(await _content_pages(reader, scope, execution, "", SIGNAL_MAX_CONTENT_PAGES)) + return { + "task": SIGNAL_TASK, + "steps": tuple(step.model_dump(mode="json") for step in steps), + } + + +def _noul_score(value: object) -> float | None: + try: + return NoulAnswer.model_validate(value).noul + except ValidationError: + return None + + +class SignalClassifier: + def __init__(self, reader: SourceReader, completion: DecisionsCall, clock: Clock) -> None: + self.reader: Final = reader + self.completion: Final = completion + self.clock: Final = clock + + async def classify(self, scope: Scope, execution: Execution, config: SignalConfig) -> SignalAttempt: + try: + state: Final = await signal_state(self.reader, scope, execution) + questions: Final = { + signal.id: {"type": "noul", "instructions": signal.question} for signal in config.signals + } + with lens_analysis(), inherit_message_logging_privacy(True): + response: Final = await self.completion( + model=config.model, + state=state, + questions=questions, + timeout=60, + metadata={"tags": ["litellm-lens-signals"]}, + ) + output: Final = DecisionsOutput.model_validate(response) + scores: Final = MappingProxyType( + { + signal.id: score + for signal in config.signals + if (score := _noul_score(output.answers.get(signal.id))) is not None + } + ) + if len(scores) != len(config.signals): + return SignalAttempt( + status="failed", + scores=scores, + model=config.model, + error="Decisions response omitted a configured noul answer", + ) + return SignalAttempt(status="classified", scores=scores, model=config.model) + except Exception as error: + detail: Final = redact_internal_details(str(error))[:300] + return SignalAttempt(status="failed", model=config.model, error=detail) + + +def trace_signals( + trace: TraceIdentity, + existing: StoredTraceSignal | None, + config: SignalConfig, +) -> TraceSignals: + if existing is None or existing.config_key != config.key(): + return TraceSignals(trace_id=trace.trace_id, trace_ref=trace.trace_ref, status="unclassified") + data: Final = SignalData.model_validate(existing.data) + if data.status == "pending": + return TraceSignals( + trace_id=trace.trace_id, + trace_ref=trace.trace_ref, + status="pending", + model=data.model, + ) + if data.status == "failed" or data.error: + return TraceSignals( + trace_id=trace.trace_id, + trace_ref=trace.trace_ref, + status="failed", + model=data.model, + classified_at=existing.classified_at, + ) + flags: Final = tuple( + sorted( + ( + SignalFlag(signal_id=signal.id, name=signal.name, score=data.scores[signal.id]) + for signal in config.signals + if signal.id in data.scores and data.scores[signal.id] >= config.threshold + ), + key=lambda flag: flag.score, + reverse=True, + ) + ) + return TraceSignals( + trace_id=trace.trace_id, + trace_ref=trace.trace_ref, + status="classified", + flags=flags, + model=data.model, + classified_at=existing.classified_at, + ) + + +async def _process_claimed( + classifier: SignalClassifier, + repository: SignalRepositoryProtocol, + scope: Scope, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, +) -> None: + from litellm._logging import verbose_proxy_logger + + attempt: Final = await classifier.classify(scope, execution, config) + try: + await repository.store(execution, config, claimed_until, classifier.clock(), attempt) + except Exception as error: + verbose_proxy_logger.error("Lens signal result could not be stored: %s", redact_internal_details(str(error))) + + +class _SignalScan: + def __init__( + self, + reader: SourceReader, + repository: SignalRepositoryProtocol, + scope: Scope, + config: SignalConfig, + now: datetime, + cursor: str, + limit: int, + ) -> None: + self.reader: Final = reader + self.repository: Final = repository + self.scope: Final = scope + self.config: Final = config + self.now: Final = now + self.cursor: str = cursor + self.limit: Final = limit + self.executions: tuple[Execution, ...] = () + self.finished: bool = False + + async def _read_page(self, start: int, end: int) -> tuple[tuple[Execution, ...], str | None]: + page_cursor: Final = self.cursor + sample: Final = await self.reader.sample( + self.scope, + ActivitySelection(source="traces"), + start, + end, + page_size=SIGNAL_PAGE_SIZE, + cursor=page_cursor, + ) + identities: Final = tuple( + TraceIdentity(trace_id=trace.trace_id, trace_ref=trace.trace_ref) for trace in sample.executions + ) + existing_rows: Final = await self.repository.traces(identities) + existing: Final = MappingProxyType({signal_identity(row): row for row in existing_rows}) + remaining: Final = self.limit - len(self.executions) + all_eligible: Final = tuple( + execution + for execution in sample.executions + if candidate(execution, existing.get(signal_identity(execution)), self.config.key(), self.now) + ) + eligible: Final = all_eligible[:remaining] + next_cursor: Final = page_cursor if len(all_eligible) > remaining else sample.next_cursor + 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): + if self.finished or len(self.executions) >= self.limit: + break + eligible, next_cursor = await self._read_page(start, end) + self.executions = (*self.executions, *eligible) + if next_cursor is None: + self.cursor = "" + self.finished = True + else: + self.cursor = next_cursor + return self.executions, self.cursor + + +async def _scan_pages( + reader: SourceReader, + repository: SignalRepositoryProtocol, + scope: Scope, + config: SignalConfig, + now: datetime, + cursor: str, + remaining: int, +) -> tuple[tuple[Execution, ...], str]: + if remaining <= 0: + return (), cursor + scan: Final = _SignalScan(reader, repository, scope, config, now, cursor, remaining) + return await scan.run() + + +async def run_signal_tick( + storage: Storage, + repository: SignalRepositoryProtocol | None, + completion: DecisionsCall | None, + clock: Clock, + router_ready: RouterReady = lambda: True, + cursor: str = "", +) -> str: + if repository is None or completion is None or not router_ready(): + return cursor + now: Final = clock() + config: Final = await repository.get_config() + if not config.enabled: + return cursor + reader: Final = SourceReader(storage) + scope: Final = Scope(all_teams=True) + candidates: Final = await _scan_pages( + reader, + repository, + scope, + config, + now, + cursor, + SIGNAL_MAX_PER_TICK, + ) + executions, next_cursor = candidates + classifier: Final = SignalClassifier(reader, completion, clock) + semaphore: Final = asyncio.Semaphore(SIGNAL_CONCURRENCY) + + async def process(execution: Execution) -> None: + from litellm._logging import verbose_proxy_logger + + async with semaphore: + claimed_at: Final = classifier.clock() + claimed_until: Final = claimed_at + SIGNAL_CLAIM_LEASE + try: + 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 + if not claimed: + return + await _process_claimed(classifier, repository, scope, execution, config, claimed_until) + + await asyncio.gather(*(process(execution) for execution in executions)) + return next_cursor + + +class _SignalLoopState: + def __init__(self) -> None: + self.cursor: str = "" + + +async def run_signal_loop( + storage: Storage, + repository: SignalRepositoryProtocol | None, + completion: DecisionsCall | None, + clock: Clock = lambda: datetime.now(timezone.utc), + router_ready: RouterReady = lambda: True, +) -> 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) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index db6294a364d..68e10467c63 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -581,6 +581,14 @@ from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_sp from litellm.proxy.image_endpoints.endpoints import router as image_router from litellm.proxy.lens.dataset_endpoints import router as lens_dataset_router 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 ( + DecisionQuestions, + DecisionsCall, + DecisionState, + run_signal_loop, +) from litellm.proxy.list_api.common import ( ManagementProblem, problem_response, @@ -1275,6 +1283,27 @@ async def _connect_to_count_stored_values() -> SupportsRawQueries: return client.writer_db +async def _call_current_lens_signal_router( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], +) -> object: + current_router: Final = llm_router + if current_router is None: + raise RuntimeError("The proxy router is not initialized") + decisions: Final[DecisionsCall] = cast(DecisionsCall, current_router.adecisions) + return await decisions( + model=model, + state=state, + questions=questions, + timeout=timeout, + metadata=metadata, + ) + + @asynccontextmanager async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]: global \ @@ -1645,12 +1674,30 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState state: Final[ProxyLifespanState] = {"tracing_receiver": receiver} 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, + ) + ) + if receiver is not None and prisma_client is not None + else None + ) + try: async with AsyncExitStack() as admin_mcp_stack: try: await admin_mcp_stack.enter_async_context(admin_mcp_lifespan(app)) yield state finally: + if signal_task is not None: + signal_task.cancel() + await asyncio.gather(signal_task, return_exceptions=True) + if model_info_scheduler is not None and model_info_scheduler.running: model_info_scheduler.remove_job("refresh_model_info") if model_info_scheduler is not scheduler: @@ -13951,7 +13998,7 @@ async def run_thread( # ) # async def get_available_routes(user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth)): from litellm.llms.base_llm.base_utils import BaseTokenCounter -from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient +from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient, writer_wrapper from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.model_repository import ModelRepository from litellm.repositories.table_repositories import ( diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index ceb127c31bd..3b83c5b09cc 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1977,3 +1977,20 @@ model LiteLLM_LensDataset { @@id([id, revision]) } + +model LiteLLM_LensSignalConfig { + id String @id + data Json +} + +model LiteLLM_LensTraceSignal { + trace_id String + trace_ref String @default("") + config_key String + span_count Int + claimed_until DateTime? + classified_at DateTime? + data Json + + @@id([trace_id, trace_ref]) +} diff --git a/schema.prisma b/schema.prisma index ceb127c31bd..3b83c5b09cc 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1977,3 +1977,20 @@ model LiteLLM_LensDataset { @@id([id, revision]) } + +model LiteLLM_LensSignalConfig { + id String @id + data Json +} + +model LiteLLM_LensTraceSignal { + trace_id String + trace_ref String @default("") + config_key String + span_count Int + claimed_until DateTime? + classified_at DateTime? + data Json + + @@id([trace_id, trace_ref]) +} diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index 6e393337524..e1c04e3d6f4 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -1,25 +1,30 @@ -from collections.abc import Callable +import asyncio +from collections.abc import Callable, Mapping from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final import pytest from fastapi import HTTPException -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError import litellm from litellm import Router from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.lens.endpoints import ( claim_due, + get_signals, list_agents, + put_signals, read_reviews, result, run_settings, run_window, trace_findings, + trace_signal_statuses, user_scope, validate_model, + validate_signal_model, watchable, watching, worker_supports_model, @@ -43,6 +48,7 @@ from litellm.proxy.lens.models import ( Worker, ) from litellm.proxy.lens.repository import DueLens, Row +from litellm.proxy.lens.signals import SignalConfig, StoredTraceSignal from litellm.proxy.lens.state import claim_job, queue_job, replace_job from litellm.rust_bridge.trace.generated.models import ExecutionRow, LensSampleParams from tests.unit.proxy.lens.test_agent_workspace import execution @@ -71,6 +77,46 @@ class ResultDatabase: return len(self.completed) +class SignalStatusDatabase: + def __init__(self, config: SignalConfig, rows: Mapping[str, StoredTraceSignal]) -> None: + self.config: Final = config + self.rows: Final = rows + self.saved: Final[asyncio.Queue[tuple[object, ...]]] = asyncio.Queue() + + async def query_raw(self, query: str, *args: object) -> object: + if '"LiteLLM_LensSignalConfig"' in query: + return ({"data": self.config.model_dump(mode="json")},) + payload: Final = args[0] + assert isinstance(payload, str) + requested: Final = TypeAdapter(tuple[TraceIdentity, ...]).validate_json(payload) + return tuple( + {"data": row.model_dump(mode="json")} + for identity in requested + if (row := self.rows.get(identity.trace_id)) is not None + ) + + async def execute_raw(self, query: str, *args: object) -> int: + await self.saved.put(args) + return 1 + + +def signal_router() -> Router: + return Router( + model_list=[ + { + "model_name": "decision", + "litellm_params": {"model": "openai/test-decision", "api_key": "test-key"}, + "model_info": {"mode": "evaluation"}, + }, + { + "model_name": "chat", + "litellm_params": {"model": "openai/test-chat", "api_key": "test-key"}, + "model_info": {"mode": "chat"}, + }, + ] + ) + + @pytest.mark.asyncio async def test_worker_sample_retries_oversized_pages_and_keeps_all_executions( monkeypatch: pytest.MonkeyPatch, @@ -469,6 +515,144 @@ async def test_trace_finding_counts_require_investigation_read_access() -> None: assert error.value.status_code == 403 +@pytest.mark.asyncio +async def test_signal_endpoints_return_statuses_in_request_order_for_admin_viewers( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy import proxy_server + + config: Final = SignalConfig(model="decision") + rows: Final = { + "pending": StoredTraceSignal( + trace_id="pending", + config_key=config.key(), + span_count=1, + claimed_until=NOW + timedelta(minutes=1), + data={"status": "pending", "scores": {}, "model": "decision", "error": ""}, + ), + "classified": StoredTraceSignal( + trace_id="classified", + config_key=config.key(), + span_count=1, + classified_at=NOW, + data={ + "status": "classified", + "scores": {"user_frustration": 0.7, "missing_capability": 0.8}, + "model": "decision", + "error": "", + }, + ), + "failed": StoredTraceSignal( + trace_id="failed", + config_key=config.key(), + span_count=1, + classified_at=NOW, + data={"status": "failed", "scores": {}, "model": "decision", "error": "classification failed"}, + ), + "stale": StoredTraceSignal( + trace_id="stale", + config_key="old-config", + span_count=1, + classified_at=NOW, + data={ + "status": "classified", + "scores": {"user_frustration": 1.0}, + "model": "old", + "error": "", + }, + ), + } + database: Final = SignalStatusDatabase(config, rows) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=database)) + viewer: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + request: Final = TraceFindingsRequest( + traces=tuple( + TraceIdentity(trace_id=trace_id) for trace_id in ("failed", "classified", "missing", "pending", "stale") + ) + ) + + assert await get_signals(viewer) == config + results: Final = await trace_signal_statuses(request, viewer) + + assert tuple((result.trace_id, result.status) for result in results) == ( + ("failed", "failed"), + ("classified", "classified"), + ("missing", "unclassified"), + ("pending", "pending"), + ("stale", "unclassified"), + ) + assert tuple((flag.signal_id, flag.name, flag.score) for flag in results[1].flags) == ( + ("missing_capability", "Missing capability", 0.8), + ("user_frustration", "User frustration", 0.7), + ) + + +@pytest.mark.asyncio +async def test_signal_endpoints_require_connected_postgres(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "prisma_client", None) + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + + with pytest.raises(HTTPException) as error: + await get_signals(auth) + + assert error.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_put_signals_saves_config_for_admin(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + database: Final = SignalStatusDatabase(SignalConfig(), {}) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=database)) + monkeypatch.setattr(proxy_server, "llm_router", signal_router()) + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + body: Final = SignalConfig(model="decision", threshold=0.7) + + assert await put_signals(body, auth) == body + + saved: Final = await database.saved.get() + assert saved[0] == "global" + assert isinstance(saved[1], str) + assert SignalConfig.model_validate_json(saved[1]) == body + + +def test_signal_model_requires_a_ready_router() -> None: + with pytest.raises(HTTPException) as error: + validate_signal_model(SignalConfig(model="decision"), None) + + assert error.value.status_code == 400 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role", (LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)) +async def test_put_signals_rejects_non_admin_roles(role: LitellmUserRoles) -> None: + auth: Final = UserAPIKeyAuth(user_role=role) + with pytest.raises(HTTPException) as error: + await put_signals(SignalConfig(model="decision"), auth) + assert error.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ("chat", "unconfigured")) +async def test_put_signals_rejects_chat_and_unknown_model_groups(monkeypatch: pytest.MonkeyPatch, model: str) -> None: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "llm_router", signal_router()) + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with pytest.raises(HTTPException) as error: + await put_signals(SignalConfig(model=model), auth) + + assert error.value.status_code == 400 + assert error.value.detail == "Choose a System 1 model (evaluation mode) configured on this proxy" + + +def test_signal_model_accepts_only_evaluation_mode_groups() -> None: + assert validate_signal_model(SignalConfig(model="decision"), signal_router()) is None + + @pytest.mark.parametrize( "role", (LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.TEAM), diff --git a/tests/unit/proxy/lens/test_signals.py b/tests/unit/proxy/lens/test_signals.py new file mode 100644 index 00000000000..d2c4eda3966 --- /dev/null +++ b/tests/unit/proxy/lens/test_signals.py @@ -0,0 +1,1002 @@ +import asyncio +import json +from collections.abc import AsyncGenerator, Mapping, Sequence +from contextlib import asynccontextmanager +from datetime import datetime, timedelta, timezone +from itertools import chain +from types import MappingProxyType, SimpleNamespace +from typing import Final + +import pytest +from pydantic import JsonValue, TypeAdapter, ValidationError + +from litellm.proxy.lens.models import Execution, Scope, TraceIdentity +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_CLAIM_LEASE, + SIGNAL_MAX_SCAN_PAGES, + SIGNAL_TASK, + DecisionQuestions, + DecisionState, + Signal, + SignalAttempt, + SignalClassifier, + SignalConfig, + SignalData, + SignalStep, + StoredTraceSignal, + candidate, + run_signal_loop, + run_signal_tick, + signal_state, + trace_signals, +) +from litellm.proxy.lens.sources import SourceReader +from litellm.rust_bridge.trace.generated.models import ( + ActivityAvailability, + AgentRow, + CountRow, + ExecutionRow, + LensAccessParams, + LensContentParams, + LensEvidenceParams, + LensSampleParams, + PartRow, +) +from litellm.types.decisions import DecisionsResponse +from litellm.types.decisions import NoulAnswer as DecisionsNoulAnswer + +NOW: Final = datetime(2026, 10, 7, 12, tzinfo=timezone.utc) +CURRENT_CONFIG_KEY: Final = SignalConfig(model="decision").key() +_SIGNAL_STEPS: Final[TypeAdapter[tuple[SignalStep, ...]]] = TypeAdapter(tuple[SignalStep, ...]) +_STORED_DATA: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) + + +def execution(identity: str, span_count: int = 1) -> Execution: + return Execution( + id=identity, + source="traces", + trace_id=identity, + team_id="", + name=identity, + start_time="", + span_count=span_count, + ) + + +def part(identity: str, content: str) -> PartRow: + return PartRow( + span_id=identity, + parent_span_id="", + name=identity, + kind="agent", + start_time="", + end_time="", + content=content, + truncated=0, + ) + + +def stored_trace( + config_key: str, + *, + trace_id: str = "trace", + status: str = "classified", + span_count: int = 1, + claimed_until: datetime | None = None, + classified_at: datetime | None = NOW - timedelta(minutes=10), + scores: dict[str, float] | None = None, + error: str = "", +) -> StoredTraceSignal: + return StoredTraceSignal( + trace_id=trace_id, + trace_ref="", + config_key=config_key, + span_count=span_count, + claimed_until=claimed_until, + classified_at=classified_at, + data=_STORED_DATA.validate_python( + { + "status": status, + "scores": scores or {}, + "model": "decision", + "error": error, + } + ), + ) + + +class SignalStorage: + def __init__( + self, + executions: tuple[ExecutionRow, ...] = (), + parts: tuple[PartRow, ...] = (), + ) -> None: + self.executions: Final = executions + self.parts: Final = parts + + async def lens_availability(self, parameters: LensAccessParams) -> Sequence[ActivityAvailability]: + return () + + async def lens_agents(self, parameters: LensAccessParams) -> Sequence[AgentRow]: + return () + + async def lens_sample(self, parameters: LensSampleParams) -> Sequence[ExecutionRow]: + return self.executions + + async def lens_content(self, parameters: LensContentParams) -> Sequence[PartRow]: + return self.parts or (part(parameters.id, parameters.id),) + + async def lens_evidence(self, parameters: LensEvidenceParams) -> Sequence[CountRow]: + return () + + +class PagedSignalStorage(SignalStorage): + def __init__(self, pages: tuple[tuple[PartRow, ...], ...]) -> None: + super().__init__() + self.pages: Final = pages + + async def lens_content(self, parameters: LensContentParams) -> Sequence[PartRow]: + index: Final = int(parameters.cursor) if parameters.cursor else 0 + return self.pages[index] + + +class PagedSampleStorage(SignalStorage): + def __init__(self, pages: tuple[tuple[ExecutionRow, ...], ...], initial_cursor: str = "") -> None: + super().__init__() + self.pages: Final = pages + self.cursors: Final[asyncio.Queue[str]] = asyncio.Queue() + self.page_by_cursor: Final = MappingProxyType( + { + initial_cursor: 0, + **{page[-1].selection_key: index + 1 for index, page in enumerate(pages[:-1])}, + } + ) + + async def lens_sample(self, parameters: LensSampleParams) -> Sequence[ExecutionRow]: + await self.cursors.put(parameters.after) + index: Final = self.page_by_cursor[parameters.after] + return self.pages[index] + + +class SignalDatabase: + def __init__( + self, + config: SignalConfig | None, + *, + stored_rows: tuple[StoredTraceSignal, ...] = (), + claim_result: bool = True, + ) -> None: + self.config: Final = config + self.stored_rows: Final = stored_rows + self.claim_result: Final = claim_result + self.calls: Final[asyncio.Queue[str]] = asyncio.Queue() + self.claims: Final[asyncio.Queue[str]] = asyncio.Queue() + self.claim_args: Final[asyncio.Queue[tuple[object, ...]]] = asyncio.Queue() + self.saved: Final[asyncio.Queue[tuple[object, ...]]] = asyncio.Queue() + + async def query_raw(self, query: str, *args: object) -> object: + if '"LiteLLM_LensSignalConfig"' in query: + return () if self.config is None else (Row(data=self.config.model_dump(mode="json")),) + if query.startswith("SELECT jsonb_build_object"): + payload: Final = args[0] + assert isinstance(payload, str) + requested: Final = TypeAdapter(tuple[TraceIdentity, ...]).validate_json(payload) + identities: Final = tuple((trace.trace_id, trace.trace_ref) for trace in requested) + return tuple( + Row(data=stored.model_dump(mode="json")) + for stored in self.stored_rows + if (stored.trace_id, stored.trace_ref) in identities + ) + if query.startswith('INSERT INTO "LiteLLM_LensTraceSignal"'): + await self.claim_args.put(args) + if not self.claim_result: + return () + trace_id: Final = args[0] + assert isinstance(trace_id, str) + await self.claims.put(trace_id) + return (Row(data={"trace_id": trace_id}),) + raise AssertionError(f"Unexpected query: {query}") + + async def execute_raw(self, query: str, *args: object) -> int: + await self.saved.put(args) + return 1 + + @asynccontextmanager + async def transaction(self) -> AsyncGenerator[Database, None]: + yield self + + +def saved_result(args: tuple[object, ...]) -> SignalData: + payload: Final = args[1] + assert isinstance(payload, str) + return SignalData.model_validate_json(payload) + + +@pytest.mark.asyncio +async def test_signal_repository_reads_defaults_and_saves_the_global_config() -> None: + database: Final = SignalDatabase(None) + repository: Final = SignalRepository(database) + updated: Final = SignalConfig(model="decision", threshold=0.7) + + assert await repository.get_config() == SignalConfig() + await repository.save_config(updated) + + saved: Final = await database.saved.get() + assert saved[0] == "global" + assert isinstance(saved[1], str) + assert SignalConfig.model_validate_json(saved[1]) == updated + + +@pytest.mark.asyncio +async def test_signal_repository_reads_rows_and_reports_a_lost_claim() -> None: + config: Final = SignalConfig(model="decision") + row: Final = stored_trace(config.key()) + database: Final = SignalDatabase(config, stored_rows=(row,), claim_result=False) + repository: Final = SignalRepository(database) + + assert await repository.traces(()) == () + assert await repository.traces((TraceIdentity(trace_id="trace"),)) == (row,) + assert not await repository.claim(execution("trace"), config, NOW + timedelta(minutes=5), NOW) + + +def test_signal_config_hashes_questions_but_not_threshold_or_display_name() -> None: + config: Final = SignalConfig(model="decision") + different_threshold: Final = config.model_copy(update={"threshold": 0.9}) + renamed: Final = config.model_copy( + update={ + "signals": ( + config.signals[0].model_copy(update={"name": "Frustration"}), + *config.signals[1:], + ) + } + ) + changed_question: Final = config.model_copy( + update={ + "signals": ( + config.signals[0].model_copy(update={"question": "Does this user sound upset?"}), + *config.signals[1:], + ) + } + ) + + assert config.key() == different_threshold.key() == renamed.key() + assert config.key() != changed_question.key() + assert DEFAULT_SIGNALS == config.signals + + +def test_signal_config_rejects_duplicate_ids_and_non_finite_thresholds() -> None: + duplicate: Final = Signal(id="same", name="First", question="Question one") + with pytest.raises(ValidationError): + SignalConfig(signals=(duplicate, duplicate)) + with pytest.raises(ValidationError): + SignalConfig(threshold=float("nan")) + + +@pytest.mark.parametrize( + "stored,trace_count,expected", + ( + (None, 1, True), + (stored_trace("old"), 1, True), + ( + stored_trace(CURRENT_CONFIG_KEY, span_count=1, classified_at=NOW - timedelta(minutes=6)), + 2, + True, + ), + ( + stored_trace( + CURRENT_CONFIG_KEY, + status="failed", + classified_at=(NOW - timedelta(minutes=31)).replace(tzinfo=None), + ), + 1, + True, + ), + (stored_trace(CURRENT_CONFIG_KEY), 1, False), + ( + stored_trace(CURRENT_CONFIG_KEY, span_count=1, classified_at=NOW - timedelta(minutes=2)), + 2, + False, + ), + ( + stored_trace("old", claimed_until=NOW + timedelta(minutes=1)), + 1, + False, + ), + ( + stored_trace( + CURRENT_CONFIG_KEY, + status="pending", + claimed_until=(NOW + timedelta(minutes=1)).replace(tzinfo=None), + classified_at=None, + ), + 1, + False, + ), + ( + stored_trace( + CURRENT_CONFIG_KEY, + status="pending", + claimed_until=(NOW - timedelta(minutes=1)).replace(tzinfo=None), + classified_at=None, + ), + 1, + True, + ), + (stored_trace(CURRENT_CONFIG_KEY, span_count=2), 1, False), + (stored_trace(CURRENT_CONFIG_KEY, span_count=1, classified_at=None), 2, False), + ), +) +def test_candidate_selection_respects_config_span_age_failure_age_and_claims( + stored: StoredTraceSignal | None, trace_count: int, expected: bool +) -> None: + config: Final = SignalConfig(model="decision") + assert candidate(execution("trace", trace_count), stored, config.key(), NOW) is expected + + +@pytest.mark.asyncio +async def test_classifier_sends_noul_questions_and_keeps_every_signal_score() -> None: + config: Final = SignalConfig(model="decision") + run: Final = execution("trace") + storage: Final = SignalStorage(parts=(part("agent", "user asks for a result"),)) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + assert model == "decision" + assert state == { + "task": SIGNAL_TASK, + "steps": ({"kind": "agent", "name": "agent", "content": "user asks for a result"},), + } + assert _STORED_DATA.validate_json(json.dumps(state)) == { + "task": SIGNAL_TASK, + "steps": [{"kind": "agent", "name": "agent", "content": "user asks for a result"}], + } + expected_questions: Final = { + signal.id: {"type": "noul", "instructions": signal.question} for signal in config.signals + } + assert questions == expected_questions + assert _STORED_DATA.validate_json(json.dumps(questions)) == expected_questions + assert timeout == 60 + assert metadata == {"tags": ["litellm-lens-signals"]} + return DecisionsResponse( + answers={ + "user_frustration": DecisionsNoulAnswer(type="noul", noul=0.9), + "missing_capability": DecisionsNoulAnswer(type="noul", noul=0.6), + "repeated_request": DecisionsNoulAnswer(type="noul", noul=0.2), + "unknown": DecisionsNoulAnswer(type="noul", noul=1.0), + } + ) + + attempt: Final = await SignalClassifier(SourceReader(storage), decide, lambda: NOW).classify( + Scope(all_teams=True), run, config + ) + + assert attempt == SignalAttempt( + status="classified", + scores={"user_frustration": 0.9, "missing_capability": 0.6, "repeated_request": 0.2}, + model="decision", + ) + + +@pytest.mark.asyncio +async def test_missing_noul_answer_fails_while_unknown_and_non_noul_answers_are_ignored() -> None: + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "choice", "choice": "yes"}, + "unknown": {"type": "noul", "noul": 1.0}, + } + } + + attempt: Final = await SignalClassifier( + SourceReader(SignalStorage(parts=(part("agent", "content"),))), + decide, + lambda: NOW, + ).classify(Scope(all_teams=True), execution("trace"), SignalConfig(model="decision")) + + assert attempt.status == "failed" + assert attempt.scores == {"user_frustration": 0.9} + assert attempt.error == "Decisions response omitted a configured noul answer" + + +@pytest.mark.asyncio +async def test_classifier_turns_decisions_errors_into_failed_attempts() -> None: + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + raise RuntimeError("decisions unavailable") + + attempt: Final = await SignalClassifier( + SourceReader(SignalStorage(parts=(part("agent", "content"),))), + decide, + lambda: NOW, + ).classify(Scope(all_teams=True), execution("trace"), SignalConfig(model="decision")) + + assert attempt == SignalAttempt(status="failed", model="decision", error="decisions unavailable") + + +def test_signal_flags_use_current_threshold_and_current_display_name() -> None: + config: Final = SignalConfig(model="decision") + row: Final = stored_trace( + config.key(), + scores={"user_frustration": 0.91, "missing_capability": 0.67, "repeated_request": 0.49}, + ) + high_threshold: Final = config.model_copy( + update={ + "threshold": 0.9, + "signals": ( + config.signals[0].model_copy(update={"name": "Frustrated user"}), + *config.signals[1:], + ), + } + ) + trace: Final = TraceIdentity(trace_id="trace") + lower: Final = trace_signals(trace, row, config) + higher: Final = trace_signals(trace, row, high_threshold) + + assert config.key() == high_threshold.key() + assert tuple((flag.signal_id, flag.score) for flag in lower.flags) == ( + ("user_frustration", 0.91), + ("missing_capability", 0.67), + ) + assert tuple((flag.signal_id, flag.name, flag.score) for flag in higher.flags) == ( + ("user_frustration", "Frustrated user", 0.91), + ) + assert not candidate(execution("trace"), row, high_threshold.key(), NOW) + + +def test_signal_flags_report_stored_errors_even_when_the_status_is_classified() -> None: + config: Final = SignalConfig(model="decision") + row: Final = stored_trace(config.key(), status="classified", error="classification failed") + + result: Final = trace_signals(TraceIdentity(trace_id="trace"), row, config) + + assert result.status == "failed" + assert result.model == "decision" + assert result.classified_at == row.classified_at + + +@pytest.mark.asyncio +async def test_signal_state_caps_content_to_head_and_tail_with_omitted_step() -> None: + parts: Final = tuple(part(str(index), chr(97 + index) * 2000) for index in range(30)) + state: Final = await signal_state( + SourceReader(SignalStorage(parts=parts)), + Scope(all_teams=True), + execution("trace"), + ) + steps_value: Final = state["steps"] + assert isinstance(steps_value, tuple) + steps: Final = _SIGNAL_STEPS.validate_python(steps_value) + head: Final = steps[:8] + marker: Final = steps[8] + tail: Final = steps[9:] + + assert state["task"] == SIGNAL_TASK + assert sum(len(step.content) for step in head) == 15000 + assert sum(len(step.content) for step in tail) == 25000 + assert head[0].content == "a" * 2000 + assert head[-1].content == "h" * 1000 + assert marker == SignalStep(kind="omitted", name="", content="9 steps omitted") + assert tail[0].content == "r" * 1000 + assert tail[-1].content == "~" * 2000 + + +@pytest.mark.asyncio +async def test_signal_state_limits_content_pages_and_part_sizes() -> None: + pages: Final = tuple( + tuple(part(f"page-{page}-{index}", "x" * 2501 if index == 0 else "x") for index in range(39)) + + (part(str(page + 1), "x"),) + for page in range(4) + ) + state: Final = await signal_state( + SourceReader(PagedSignalStorage(pages)), + Scope(all_teams=True), + execution("trace"), + ) + steps: Final = _SIGNAL_STEPS.validate_python(state["steps"]) + + assert len(steps) == 120 + assert steps[0].content.startswith("x" * 800) + assert "[... 501 characters omitted ...]" in steps[0].content + assert steps[0].content.endswith("x" * 1200) + assert steps[-1].name == "3" + assert all(not step.name.startswith("page-3-") for step in steps) + + small_state: Final = await signal_state( + SourceReader(SignalStorage(parts=(part("small", "ok"),))), + Scope(all_teams=True), + execution("trace"), + ) + small_steps: Final = _SIGNAL_STEPS.validate_python(small_state["steps"]) + assert small_steps == (SignalStep(kind="agent", name="small", content="ok"),) + + +@pytest.mark.asyncio +async def test_signal_state_part_excerpt_preserves_the_output_tail() -> None: + content: Final = "I" * 5000 + "OUTPUT: refused" + state: Final = await signal_state( + SourceReader(SignalStorage(parts=(part("result", content),))), + Scope(all_teams=True), + execution("trace"), + ) + steps: Final = _SIGNAL_STEPS.validate_python(state["steps"]) + excerpt: Final = steps[0].content + marker: Final = "\n[... 3015 characters omitted ...]\n" + + assert marker in excerpt + assert excerpt.endswith("OUTPUT: refused") + assert len(excerpt) == 800 + len(marker) + 1200 + + +@pytest.mark.asyncio +async def test_signal_tick_classifies_at_most_50_traces_and_persists_scores() -> None: + config: Final = SignalConfig(model="decision") + executions: Final = tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{index}", + team_id="", + name=f"trace-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=60, + selected=60, + selection_key=f"cursor-{index}", + ) + for index in range(60) + ) + storage: Final = SignalStorage(executions=executions) + database: Final = SignalDatabase(config) + repository: Final = SignalRepository(database) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + steps: Final = TypeAdapter(tuple[SignalStep, ...]).validate_python(state["steps"]) + await database.calls.put(steps[0].name) + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "noul", "noul": 0.6}, + "repeated_request": {"type": "noul", "noul": 0.2}, + } + } + + await run_signal_tick(storage, repository, decide, lambda: NOW) + classified: Final = tuple(database.saved.get_nowait() for _ in range(database.saved.qsize())) + traces: Final = tuple(database.calls.get_nowait() for _ in range(database.calls.qsize())) + saved_data: Final = tuple(saved_result(args) for args in classified) + + assert len(classified) == 50 + assert frozenset(traces) == frozenset(f"trace-{index}" for index in range(50)) + assert ( + saved_data + == ( + SignalData( + status="classified", + scores={ + "user_frustration": 0.9, + "missing_capability": 0.6, + "repeated_request": 0.2, + }, + model="decision", + error="", + ), + ) + * 50 + ) + + +@pytest.mark.asyncio +async def test_signal_tick_claims_with_worker_start_time_and_skips_lost_claims() -> None: + config: Final = SignalConfig(model="decision") + executions: Final = tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{index}", + team_id="", + name=f"trace-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=2, + selected=2, + selection_key=f"cursor-{index}", + ) + for index in range(2) + ) + database: Final = SignalDatabase(config, claim_result=False) + repository: Final = SignalRepository(database) + + class AdvancingClock: + def __init__(self) -> None: + self.values: Final = tuple(NOW + timedelta(minutes=index) for index in range(3)) + self.index: int = 0 + + def __call__(self) -> datetime: + value: Final = self.values[self.index] + self.index += 1 + return value + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + await run_signal_tick(SignalStorage(executions=executions), repository, decide, AdvancingClock()) + + claims: Final = tuple(database.claim_args.get_nowait() for _ in range(database.claim_args.qsize())) + + def claim_times(args: tuple[object, ...]) -> tuple[datetime, datetime]: + claimed_until: Final = args[4] + claimed_at: Final = args[6] + assert isinstance(claimed_until, datetime) + assert isinstance(claimed_at, datetime) + return claimed_until, claimed_at + + times: Final = tuple(claim_times(claim) for claim in claims) + assert database.calls.empty() + assert database.saved.empty() + assert all(claimed_until == claimed_at + SIGNAL_CLAIM_LEASE for claimed_until, claimed_at in times) + assert all(claimed_at != NOW for _, claimed_at in times) + + +@pytest.mark.asyncio +async def test_signal_tick_resumes_after_ten_pages_and_resets_after_a_short_page() -> None: + config: Final = SignalConfig(model="decision") + + def sample_page(page: int) -> tuple[ExecutionRow, ...]: + return tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{page}-{index}", + team_id="", + name=f"trace-{page}-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=2500, + selected=2500, + selection_key=f"page-{page}-{index}", + ) + for index in range(100) + ) + + pages: Final = tuple(sample_page(page) for page in range(25)) + all_rows: Final = tuple(chain.from_iterable(pages)) + stored_rows: Final = tuple(stored_trace(CURRENT_CONFIG_KEY, trace_id=row.trace_id) for row in all_rows) + storage: Final = PagedSampleStorage(pages) + database: Final = SignalDatabase(config, stored_rows=stored_rows) + repository: Final = SignalRepository(database) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + first_cursor: Final = await run_signal_tick(storage, repository, decide, lambda: NOW) + 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_calls: Final = tuple(storage.cursors.get_nowait() for _ in range(storage.cursors.qsize())) + + assert len(first_calls) == SIGNAL_MAX_SCAN_PAGES + assert first_cursor + assert len(second_calls) == SIGNAL_MAX_SCAN_PAGES + assert second_calls[0] == first_cursor + assert second_cursor + + 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, + ) + assert short_cursor == "" + + +@pytest.mark.asyncio +async def test_signal_tick_resumes_a_partially_consumed_page() -> None: + config: Final = SignalConfig(model="decision") + page: Final = tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{index}", + team_id="", + name=f"trace-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=100, + selected=100, + selection_key=f"cursor-{index}", + ) + for index in range(100) + ) + initial_rows: Final = tuple(stored_trace(CURRENT_CONFIG_KEY, trace_id=f"trace-{index}") for index in range(20)) + resume_cursor: Final = "resume-page" + storage: Final = PagedSampleStorage((page, ()), initial_cursor=resume_cursor) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "noul", "noul": 0.6}, + "repeated_request": {"type": "noul", "noul": 0.2}, + } + } + + 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_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_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)) + + assert first_cursor == resume_cursor + assert second_cursor == "" + assert len(first_claims) == 50 + assert len(second_claims) == 30 + assert frozenset(first_claims).isdisjoint(second_claims) + assert frozenset(first_claims) | frozenset(second_claims) == expected_eligible + assert sample_cursors == (resume_cursor, resume_cursor, page[-1].selection_key) + + +@pytest.mark.asyncio +async def test_signal_tick_skips_claims_and_writes_when_router_is_not_ready() -> None: + config: Final = SignalConfig(model="decision") + storage: Final = SignalStorage( + executions=( + ExecutionRow( + source="traces", + trace_id="trace", + team_id="", + name="trace", + start_time="", + span_count=1, + root_seen=1, + eligible=1, + selected=1, + selection_key="cursor", + ), + ) + ) + database: Final = SignalDatabase(config) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + await run_signal_tick( + storage, + SignalRepository(database), + decide, + lambda: NOW, + router_ready=lambda: False, + ) + + assert database.claims.empty() + assert database.saved.empty() + + +@pytest.mark.asyncio +async def test_signal_tick_skips_missing_dependencies_and_disabled_configs() -> None: + storage: Final = SignalStorage() + + await run_signal_tick(storage, None, None, lambda: NOW) + + database: Final = SignalDatabase(SignalConfig()) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + raise AssertionError("disabled signal config should not call Decisions") + + await run_signal_tick(storage, SignalRepository(database), decide, lambda: NOW) + assert database.claims.empty() + assert database.saved.empty() + + +class FailingStoreDatabase(SignalDatabase): + async def execute_raw(self, query: str, *args: object) -> int: + raise RuntimeError("store unavailable") + + +@pytest.mark.asyncio +async def test_signal_tick_continues_when_storing_a_result_fails() -> None: + config: Final = SignalConfig(model="decision") + execution_row: Final = ExecutionRow( + source="traces", + trace_id="trace", + team_id="", + name="trace", + start_time="", + span_count=1, + root_seen=1, + eligible=1, + selected=1, + selection_key="cursor", + ) + database: Final = FailingStoreDatabase(config) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "noul", "noul": 0.6}, + "repeated_request": {"type": "noul", "noul": 0.2}, + } + } + + await run_signal_tick( + SignalStorage(executions=(execution_row,)), + SignalRepository(database), + decide, + lambda: NOW, + ) + + assert await database.claims.get() == "trace" + assert database.saved.empty() + + +class FailingSignalRepository: + def __init__(self) -> None: + self.started: Final = asyncio.Event() + + async def get_config(self) -> SignalConfig: + self.started.set() + await asyncio.sleep(0) + raise RuntimeError("tick failed") + + +@pytest.mark.asyncio +async def test_signal_loop_continues_after_a_tick_error() -> None: + repository: Final = FailingSignalRepository() + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + task: Final = asyncio.create_task(run_signal_loop(SignalStorage(), repository, decide, lambda: NOW)) + await repository.started.wait() + await asyncio.sleep(0) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + +@pytest.mark.asyncio +async def test_proxy_signal_call_resolves_the_current_router(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + async def first_decisions( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return "first" + + async def second_decisions( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return "second" + + async def call_current_router() -> object: + return await proxy_server._call_current_lens_signal_router( + model="decision", + state={"task": "task"}, + questions={}, + timeout=60, + metadata={"tags": ["test"]}, + ) + + monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(adecisions=first_decisions)) + assert await call_current_router() == "first" + + monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(adecisions=second_decisions)) + assert await call_current_router() == "second" + + monkeypatch.setattr(proxy_server, "llm_router", None) + with pytest.raises(RuntimeError, match="router is not initialized"): + await call_current_router() diff --git a/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx index 262d7ff9a51..f661c965e0b 100644 --- a/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx @@ -291,7 +291,7 @@ describe("Lens interactive demo", () => { await expectUrl(onUrlUpdate, (url) => expect(url.get("tab")).toBe("settings")); expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); const panel = within(screen.getByRole("region", { name: "Settings" })); - expect(panel.getByRole("status")).toHaveTextContent("Tracing enabled"); + expect(panel.getByText("Tracing enabled", { exact: true })).toBeVisible(); expect(panel.getByRole("heading", { name: "Analysis worker" })).toBeVisible(); expect(panel.getByRole("heading", { name: worker.name })).toBeVisible(); expect(panel.getByText("Connected")).toBeVisible(); diff --git a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx index 1a483f55e90..6c5629d2a07 100644 --- a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx @@ -196,6 +196,7 @@ function LensContent({ userRole, readOnly }: Omit readOnly={readOnly} canMintTracingKey={isAdmin} canViewFindings={canViewInvestigations} + onSetUpSignals={canConfigure ? showSettings : undefined} /> diff --git a/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts b/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts index 422601c44c2..ddc7019eadb 100644 --- a/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts +++ b/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts @@ -37,6 +37,8 @@ function demoLensApi(data: LensDemoData): LensApi { saveLens: readOnly, startRun: readOnly, watchAll: async () => ({ watching: [], skipped: [] }), + signalConfig: async () => ({ model: "", threshold: 0.5, signals: [] }), + saveSignalConfig: readOnly, cancelRun: readOnly, reviewFinding: readOnly, registerWorker: readOnly, @@ -90,6 +92,8 @@ function demoTracesApi(data: LensDemoData): TracesApi { ); return { ...trace, finding_count: assessed.length ? findings.size : null }; }), + signals: async (traces) => + traces.map((trace) => ({ ...trace, status: "unclassified" as const, flags: [], model: "", classified_at: null })), anyRecorded: async () => data.runs.length > 0, trace: (traceId) => found(run(traceId)?.trace), span: (traceId, spanId) => found(run(traceId)?.details.find((span) => span.span_id === spanId)), diff --git a/ui/litellm-dashboard/src/components/lens/data/queries.ts b/ui/litellm-dashboard/src/components/lens/data/queries.ts index a918d33107d..63c311ecd1c 100644 --- a/ui/litellm-dashboard/src/components/lens/data/queries.ts +++ b/ui/litellm-dashboard/src/components/lens/data/queries.ts @@ -25,6 +25,7 @@ export const lensKeys = { models: (scope: string) => [...lensKeys.all, "models", { scope }] as const, modelDetails: (scope: string) => [...lensKeys.all, "model-details", { scope }] as const, activity: (scope: string) => [...lensKeys.all, "activity-available", { scope }] as const, + signalConfig: (scope: string) => [...lensKeys.all, "signal-config", { scope }] as const, discoveries: () => [...lensKeys.all, "discovery"] as const, discovery: (scope: string, source: Settings["source"], hours: number | undefined) => [...lensKeys.discoveries(), { scope, source, hours }] as const, @@ -46,6 +47,13 @@ export const lensQueries = { modelDetails(api: LensApi) { return queryOptions({ queryKey: lensKeys.modelDetails(api.scope), queryFn: () => api.modelDetails() }); }, + signalConfig(api: LensApi) { + return queryOptions({ + queryKey: lensKeys.signalConfig(api.scope), + queryFn: () => api.signalConfig(), + staleTime: 5000, + }); + }, activity(api: LensApi, loaded: boolean) { const options = { queryKey: lensKeys.activity(api.scope), diff --git a/ui/litellm-dashboard/src/components/lens/data/service.ts b/ui/litellm-dashboard/src/components/lens/data/service.ts index 9bd7919308b..17580d3fdc7 100644 --- a/ui/litellm-dashboard/src/components/lens/data/service.ts +++ b/ui/litellm-dashboard/src/components/lens/data/service.ts @@ -13,6 +13,7 @@ import type { RunWindow, Sample, Settings, + SignalConfig, WorkerCreated, } from "../model/types"; @@ -62,6 +63,8 @@ export interface LensApi { saveLens(id: string | undefined, settings: Settings): Promise; startRun(lensId: string, request?: RunWindow): Promise; watchAll(): Promise; + signalConfig(): Promise; + saveSignalConfig(config: SignalConfig): Promise; cancelRun(lensId: string): Promise; reviewFinding(lensId: string, findingId: string, status: FindingStatus, reason: string): Promise; registerWorker(analysisKeyId: string): Promise; @@ -168,6 +171,8 @@ export function liveLensApi(client: LensClient, apiClient: ApiClient, accessToke ), startRun: (lensId, request = {}) => sent(client.POST("/lens/{lens_id}/runs", { ...lens(lensId), body: request })), watchAll: () => required(client.POST("/lens/watch-all", { headers })), + signalConfig: () => required(client.GET("/lens/signals", { headers })), + saveSignalConfig: (config) => required(client.PUT("/lens/signals", { headers, body: config })), cancelRun: (lensId) => sent(client.POST("/lens/{lens_id}/cancel", lens(lensId))), reviewFinding: (lensId, findingId, status, reason) => sent( diff --git a/ui/litellm-dashboard/src/components/lens/model/signals.ts b/ui/litellm-dashboard/src/components/lens/model/signals.ts new file mode 100644 index 00000000000..4c81ec8f3a4 --- /dev/null +++ b/ui/litellm-dashboard/src/components/lens/model/signals.ts @@ -0,0 +1,74 @@ +import type { AnalysisModelInfo, SignalConfig } from "./types"; + +export const SYSTEM_ONE_MODE = "evaluation"; + +export const signalsConfigured = (config: SignalConfig): boolean => + Boolean(config.model) && (config.signals?.length ?? 0) > 0; + +export const systemOneModels = (details: readonly AnalysisModelInfo[]): AnalysisModelInfo[] => + details + .filter((info) => info.mode === SYSTEM_ONE_MODE) + .toSorted((a, b) => a.model_group.localeCompare(b.model_group)); + +export interface LibrarySignal { + readonly id: string; + readonly name: string; + readonly summary: string; + readonly question: string; +} + +export const SIGNAL_LIBRARY: readonly LibrarySignal[] = [ + { + id: "user_frustration", + name: "User frustration", + summary: "annoyed, complaining or giving up", + question: + "Does the user show frustration, annoyance or dissatisfaction with the agent in this run, for example complaints, irritated corrections, all caps, profanity, or giving up on the task?", + }, + { + id: "missing_capability", + name: "Missing capability", + summary: "asked for something the agent can't do", + question: + "Does the user ask for something the agent cannot do in this run, so that the agent refuses, says it lacks a tool, permission, integration or data source, or fails because the capability does not exist?", + }, + { + id: "repeated_request", + name: "Repeated request", + summary: "had to ask for the same thing again", + question: + "Does the user ask for the same thing more than once in this run, usually because the agent did not deliver it the first time?", + }, + { + id: "asked_for_human", + name: "Asked for a human", + summary: "wants a person, not the agent", + question: "Does the user ask to talk to a human, a support person or a manager instead of the agent in this run?", + }, + { + id: "tool_failure", + name: "Tool failure", + summary: "a tool or step failed and stayed broken", + question: + "Does a tool call or step fail in this run with an error, exception or timeout that the agent does not recover from?", + }, + { + id: "refused_request", + name: "Refused request", + summary: "declined a reasonable ask", + question: "Does the agent refuse or decline a reasonable, allowed user request in this run?", + }, + { + id: "made_up_answer", + name: "Made-up answer", + summary: "facts or links with no source", + question: + "Does the agent state facts, numbers, links or tool results in this run that are not supported by the conversation or by any tool output?", + }, + { + id: "task_abandoned", + name: "Task abandoned", + summary: "run ended without what was asked", + question: "Does the run end without the user getting what they asked for?", + }, +]; diff --git a/ui/litellm-dashboard/src/components/lens/model/types.ts b/ui/litellm-dashboard/src/components/lens/model/types.ts index 943ca1ad5e6..2f8f54f65c9 100644 --- a/ui/litellm-dashboard/src/components/lens/model/types.ts +++ b/ui/litellm-dashboard/src/components/lens/model/types.ts @@ -5,6 +5,8 @@ export type Lens = components["schemas"]["Lens"]; export type Settings = components["schemas"]["LensSettings"]; export type LensList = components["schemas"]["LensList"]; +export type SignalConfig = components["schemas"]["SignalConfig"]; +export type Signal = NonNullable[number]; export type Finding = components["schemas"]["Finding"]; diff --git a/ui/litellm-dashboard/src/components/lens/settings/LensSettings.tsx b/ui/litellm-dashboard/src/components/lens/settings/LensSettings.tsx index ce07e12795c..92a05cf68e9 100644 --- a/ui/litellm-dashboard/src/components/lens/settings/LensSettings.tsx +++ b/ui/litellm-dashboard/src/components/lens/settings/LensSettings.tsx @@ -5,6 +5,7 @@ import { Activity, ArrowUpRight } from "lucide-react"; import { Button } from "@/components/ui/button"; import { StatusDot } from "@/components/shared/StatusDot"; import { WorkerSettings } from "./worker/WorkerSettings"; +import { SignalSettings } from "./signals/SignalSettings"; import { SettingsCard, SettingsSection } from "./SettingsSection"; import type { LensList } from "../model/types"; @@ -53,6 +54,7 @@ export function LensSettings({ return (
+ (); + +describe("signal settings", () => { + beforeEach(() => { + testQueryClient.clear(); + network.mockReset(); + network.mockImplementation(async (input) => { + const path = new URL(input instanceof Request ? input.url : String(input), "http://localhost").pathname; + return Response.json(path === "/model_group/info" ? { data: [] } : {}); + }); + vi.stubGlobal("fetch", network); + }); + + it("updates a clean draft when saved settings change", async () => { + const { rerender } = renderWithLens(); + const model = await screen.findByRole("combobox", { name: "System 1 model" }); + + expect(model).toHaveValue("jev"); + rerender(); + + expect(model).toHaveValue("new-jev"); + expect(screen.queryByRole("alert")).not.toBeInTheDocument(); + }); + + it("keeps dirty edits until the user loads the latest settings", async () => { + const user = userEvent.setup(); + const { rerender } = renderWithLens(); + const threshold = await screen.findByRole("spinbutton", { name: "Flag at score" }); + + fireEvent.change(threshold, { target: { value: "70" } }); + rerender(); + + expect(threshold).toHaveValue(70); + expect(screen.getByRole("alert")).toHaveTextContent("Signals were changed elsewhere"); + await user.click(screen.getByRole("button", { name: "Load latest" })); + + expect(threshold).toHaveValue(80); + expect(screen.queryByRole("alert")).not.toBeInTheDocument(); + }); + + it("keeps a custom Tool failure signal separate while toggling the library signal", async () => { + const user = userEvent.setup(); + const customQuestion = "Does this custom failure condition apply?"; + const customSignal = { id: "tool_failure", name: "Tool failure", question: customQuestion }; + const customConfig: SignalConfig = { + ...saved, + signals: [customSignal], + }; + renderWithLens(); + + const question = await screen.findByRole("textbox", { name: "Question for Tool failure" }); + const model = await screen.findByRole("combobox", { name: "System 1 model" }); + const threshold = screen.getByRole("spinbutton", { name: "Flag at score" }); + expect(question).toHaveValue(customQuestion); + expect(screen.getByRole("button", { name: /^Tool failure/ })).toHaveAttribute("aria-pressed", "false"); + expect(model).toHaveValue("jev"); + expect(threshold).toHaveValue(50); + + await user.click(screen.getByRole("button", { name: /^Tool failure/ })); + expect(screen.getByRole("button", { name: /^Tool failure/ })).toHaveAttribute("aria-pressed", "true"); + expect(question).toHaveValue(customQuestion); + + const editedQuestion = "Does this edited custom failure condition apply?"; + fireEvent.change(question, { target: { value: editedQuestion } }); + expect(question).toHaveValue(editedQuestion); + expect(screen.getByRole("button", { name: /^Tool failure/ })).toHaveAttribute("aria-pressed", "true"); + + await user.click(screen.getByRole("button", { name: "Remove Tool failure" })); + expect(screen.queryByRole("textbox", { name: "Question for Tool failure" })).not.toBeInTheDocument(); + expect(screen.getByRole("button", { name: /^Tool failure/ })).toHaveAttribute("aria-pressed", "true"); + + await user.click(screen.getByRole("button", { name: /^Tool failure/ })); + expect(screen.getByRole("button", { name: /^Tool failure/ })).toHaveAttribute("aria-pressed", "false"); + expect(screen.getByText("Pick at least one signal to flag traces")).toBeInTheDocument(); + expect(model).toHaveValue("jev"); + expect(threshold).toHaveValue(50); + }); + + it("removes only the library signal when toggling it off beside a custom signal", async () => { + const user = userEvent.setup(); + const customQuestion = "Does this custom failure condition apply?"; + const customConfig: SignalConfig = { + ...saved, + signals: [{ id: "tool_failure", name: "Tool failure", question: customQuestion }], + }; + renderWithLens(); + + const question = await screen.findByRole("textbox", { name: "Question for Tool failure" }); + const tile = screen.getByRole("button", { name: /^Tool failure/ }); + + await user.click(tile); + expect(tile).toHaveAttribute("aria-pressed", "true"); + await user.click(tile); + + expect(screen.getByRole("button", { name: /^Tool failure/ })).toHaveAttribute("aria-pressed", "false"); + expect(question).toHaveValue(customQuestion); + }); +}); diff --git a/ui/litellm-dashboard/src/components/lens/settings/signals/SignalSettings.tsx b/ui/litellm-dashboard/src/components/lens/settings/signals/SignalSettings.tsx new file mode 100644 index 00000000000..d80a2c768fc --- /dev/null +++ b/ui/litellm-dashboard/src/components/lens/settings/signals/SignalSettings.tsx @@ -0,0 +1,296 @@ +"use client"; + +import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query"; +import { Flag, Trash2 } from "lucide-react"; +import Link from "next/link"; +import { useId, useState } from "react"; + +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { StatusDot } from "@/components/shared/StatusDot"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Skeleton } from "@/components/ui/skeleton"; +import { Textarea } from "@/components/ui/textarea"; +import { uiHref } from "@/utils/uiHref"; + +import { useLensApi } from "../../data/LensServices"; +import { lensKeys, lensQueries } from "../../data/queries"; +import { SIGNAL_LIBRARY, signalsConfigured, systemOneModels } from "../../model/signals"; +import { WatchPicker } from "../../setup/WatchPicker"; +import type { SignalConfig } from "../../model/types"; +import { SettingsCard, SettingsSection } from "../SettingsSection"; +import { + MAX_SIGNALS, + configFrom, + draftFrom, + draftProblems, + newRow, + type SignalDraft, + type SignalRow, +} from "./signalDraft"; + +const LIBRARY_QUESTIONS: ReadonlySet = new Set(SIGNAL_LIBRARY.map((signal) => signal.question)); + +export function SignalSettings() { + return ( + + + + ); +} + +function SignalConfigLoader() { + const api = useLensApi(); + const config = useQuery(lensQueries.signalConfig(api)); + if (config.error) + return ( +

+ Could not load signals: {config.error.message} +

+ ); + if (!config.data) return ; + return ; +} + +const FieldError = ({ children }: { children?: string }) => + children ?

{children}

: null; + +function SetupCallout({ hasModels }: { hasModels: boolean }) { + return ( +
+
+ ); +} + +function SignalFields({ + row, + problems, + onChange, + onRemove, +}: { + row: SignalRow; + problems?: { name?: string; question?: string }; + onChange: (row: SignalRow) => void; + onRemove: () => void; +}) { + const label = row.name.trim() || "new signal"; + return ( +
  • +
    + onChange({ ...row, name: event.target.value })} + /> + {problems?.name} +
    +
    +