diff --git a/.github/assets/lens-datasets/add-to-dataset-dialog.png b/.github/assets/lens-datasets/add-to-dataset-dialog.png new file mode 100644 index 00000000000..8590ad48dd5 Binary files /dev/null and b/.github/assets/lens-datasets/add-to-dataset-dialog.png differ diff --git a/.github/assets/lens-datasets/dataset-case-panel.png b/.github/assets/lens-datasets/dataset-case-panel.png new file mode 100644 index 00000000000..dd579ac58af Binary files /dev/null and b/.github/assets/lens-datasets/dataset-case-panel.png differ diff --git a/.github/assets/lens-datasets/dataset-detail.png b/.github/assets/lens-datasets/dataset-detail.png new file mode 100644 index 00000000000..c68b9b6f73c Binary files /dev/null and b/.github/assets/lens-datasets/dataset-detail.png differ diff --git a/.github/assets/lens-datasets/datasets-list.png b/.github/assets/lens-datasets/datasets-list.png new file mode 100644 index 00000000000..14a6200aae0 Binary files /dev/null and b/.github/assets/lens-datasets/datasets-list.png differ diff --git a/.github/assets/lens-datasets/old-revision.png b/.github/assets/lens-datasets/old-revision.png new file mode 100644 index 00000000000..d6e369e2485 Binary files /dev/null and b/.github/assets/lens-datasets/old-revision.png differ diff --git a/.github/assets/lens-datasets/trace-header.png b/.github/assets/lens-datasets/trace-header.png new file mode 100644 index 00000000000..85bd759b5d0 Binary files /dev/null and b/.github/assets/lens-datasets/trace-header.png differ diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261006000100_lens_datasets/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261006000100_lens_datasets/migration.sql new file mode 100644 index 00000000000..80ce096d409 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261006000100_lens_datasets/migration.sql @@ -0,0 +1,7 @@ +CREATE TABLE IF NOT EXISTS "LiteLLM_LensDataset" ( + "id" TEXT NOT NULL, + "revision" INTEGER NOT NULL, + "created_at" TIMESTAMP(3) NOT NULL, + "data" JSONB NOT NULL, + CONSTRAINT "LiteLLM_LensDataset_pkey" PRIMARY KEY ("id", "revision") +); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 4aabd56131b..514df905866 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1965,3 +1965,12 @@ model LiteLLM_LensWorker { token_hash String @unique data Json } + +model LiteLLM_LensDataset { + id String + revision Int + created_at DateTime + data Json + + @@id([id, revision]) +} diff --git a/litellm-rust/crates/traces/src/ui.rs b/litellm-rust/crates/traces/src/ui.rs index eebeaab70ce..ef849d2b440 100644 --- a/litellm-rust/crates/traces/src/ui.rs +++ b/litellm-rust/crates/traces/src/ui.rs @@ -85,6 +85,32 @@ struct RawMessage { kwargs: Option>, } +#[derive(Deserialize)] +#[serde(deny_unknown_fields)] +struct AssistantSummary { + content: Option, + tool_names: Vec, +} + +impl AssistantSummary { + fn into_ui(self) -> UiMessage { + let calls: Vec = self + .tool_names + .into_iter() + .map(|name| UiToolCall { + name, + arguments: "{}".to_owned(), + }) + .collect(); + UiMessage { + role: ChatRole::Assistant, + content: self.content.unwrap_or_default(), + name: None, + tool_calls: (!calls.is_empty()).then_some(calls), + } + } +} + #[derive(Deserialize)] struct ContentBlock { #[serde(rename = "type", default)] @@ -211,6 +237,19 @@ fn messages(parsed: &Value) -> Option> { Some(unwrapped.into_iter().map(RawMessage::into_ui).collect()) } +fn assistant_summaries(parsed: &Value) -> Option> { + let summaries = Vec::::deserialize(parsed).ok()?; + if summaries.is_empty() { + return None; + } + Some( + summaries + .into_iter() + .map(AssistantSummary::into_ui) + .collect(), + ) +} + pub fn to_ui_content(raw: &str) -> UiContent { let text = || UiContent::Text { text: raw.to_owned(), @@ -223,7 +262,7 @@ pub fn to_ui_content(raw: &str) -> UiContent { Ok(parsed @ (Value::Array(_) | Value::Object(_))) => parsed, _ => return text(), }; - if let Some(messages) = messages(&parsed) { + if let Some(messages) = messages(&parsed).or_else(|| assistant_summaries(&parsed)) { return UiContent::Messages { messages }; } match parsed { @@ -356,6 +395,33 @@ mod tests { ); } + #[rstest] + fn assistant_summaries_become_assistant_messages() { + let raw = json!([ + {"content": "`/etc/hosts` has 11 lines.", "tool_names": []}, + {"content": null, "tool_names": ["terminal"]}, + ]); + assert_eq!( + to_ui_content(&raw.to_string()), + UiContent::Messages { + messages: vec![ + message("assistant", "`/etc/hosts` has 11 lines."), + UiMessage { + tool_calls: Some(vec![call("terminal", "{}")]), + ..message("assistant", "") + }, + ] + } + ); + } + + #[rstest] + #[case::extra_field(r#"[{"content": "x", "tool_names": [], "score": 1}]"#)] + #[case::missing_tool_names(r#"[{"content": "x"}]"#)] + fn near_summaries_stay_text(#[case] raw: &str) { + assert!(matches!(to_ui_content(raw), UiContent::Text { .. })); + } + #[rstest] fn plain_objects_become_fields_in_key_order() { let raw = r#"{"zeta": "plain", "alpha": {"nested": [1, 2]}, "count": 3, "missing": null}"#; diff --git a/litellm/constants.py b/litellm/constants.py index 233a0e2e4df..95d05c6d819 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -60,6 +60,9 @@ TRACE_READ_RETRY_AFTER_SECONDS: Final = get_env_int("TRACE_READ_RETRY_AFTER_SECO OTLP_MAX_CONCURRENT_INGESTS: Final = get_env_int("OTLP_MAX_CONCURRENT_INGESTS", 2) AGENT_TRACING_INPUT_PREVIEW_CHARS: Final = get_env_int("AGENT_TRACING_INPUT_PREVIEW_CHARS", 240) AGENT_TRACING_LIST_PAGE_SIZE: Final = get_env_int("AGENT_TRACING_LIST_PAGE_SIZE", 50) +LENS_DATASET_MAX_CASES: Final = get_env_int("LENS_DATASET_MAX_CASES", 200) +LENS_DATASET_MAX_CASE_CHARS: Final = get_env_int("LENS_DATASET_MAX_CASE_CHARS", 20_000) +LENS_DATASET_TRACE_PAGE_SIZE: Final = 500 DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10)) DEFAULT_S3_BATCH_SIZE: Final = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512)) DEFAULT_S3_MAX_CONCURRENT_UPLOADS: Final = int(os.getenv("DEFAULT_S3_MAX_CONCURRENT_UPLOADS", "16")) diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index 3644ba60dd5..3b17aea3cd8 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -349,6 +349,7 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset( "LiteLLM_WorkflowEvent", "LiteLLM_WorkflowMessage", "LiteLLM_Lens", + "LiteLLM_LensDataset", "LiteLLM_LensRun", "LiteLLM_LensReview", "LiteLLM_LensWorker", diff --git a/litellm/proxy/lens/dataset_endpoints.py b/litellm/proxy/lens/dataset_endpoints.py new file mode 100644 index 00000000000..1fc1019ab76 --- /dev/null +++ b/litellm/proxy/lens/dataset_endpoints.py @@ -0,0 +1,175 @@ +from datetime import datetime, timezone +from types import MappingProxyType +from typing import Annotated, Final, TypeAlias +from uuid import uuid4 + +from fastapi import APIRouter, Depends, HTTPException, Query, Response + +from litellm.constants import LENS_DATASET_TRACE_PAGE_SIZE +from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper +from litellm.proxy.lens.dataset_repository import DatasetRepository, DatasetStore +from litellm.proxy.lens.datasets import build_cases, export_jsonl, included_cases, revision_cases, revision_problem +from litellm.proxy.lens.endpoints import Auth, repository, user_scope +from litellm.proxy.lens.models import ( + BuildRequest, + BuildResult, + Dataset, + DatasetCase, + DatasetCreate, + DatasetSummary, + EvalCases, + Finding, + RevisionSave, + Scope, +) +from litellm.proxy.lens.repository import LensRepository, WriterDatabase +from litellm.proxy.lens.state import can_access +from litellm.proxy.tracing_endpoints import read_failure +from litellm.proxy.tracing_runtime import provide_receiver, require_receiver +from litellm.rust_bridge.trace.errors import TraceChanged +from litellm.rust_bridge.trace.generated.types import SpanDetail, Trace, TraceScope +from litellm.tracing import TraceReceiver + +router: Final = APIRouter(prefix="/lens/datasets", tags=["Lens"]) +ALL_TRACES: Final = TraceScope(all_teams=1, user_id="", team_ids=()) + + +def dataset_store() -> DatasetStore: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(503, "Lens needs a connected Postgres database") + return DatasetRepository(WriterDatabase(writer_wrapper(prisma_client.db))) + + +Datasets: TypeAlias = Annotated[DatasetStore, Depends(dataset_store)] +Lenses: TypeAlias = Annotated[LensRepository, Depends(repository)] +Receiver: TypeAlias = Annotated[TraceReceiver | None, Depends(provide_receiver)] +RevisionQuery: TypeAlias = Annotated[int | None, Query(ge=0)] + + +class ProxyDatasetReader: + def __init__(self, receiver: TraceReceiver | None, lenses: LensRepository, scope: Scope) -> None: + self.receiver: Final = receiver + self.lenses: Final = lenses + self.scope: Final = scope + + async def trace(self, trace_id: str, trace_ref: str) -> Trace | None: + try: + return await self._all_pages(trace_id, trace_ref) + except (TraceChanged, ValueError, OverflowError, RuntimeError) as error: + raise read_failure(error) from error + + async def _all_pages(self, trace_id: str, trace_ref: str) -> Trace | None: + receiver: Final = require_receiver(self.receiver) + first: Final = await receiver.get_trace(trace_id, ALL_TRACES, trace_ref, page_size=LENS_DATASET_TRACE_PAGE_SIZE) + if first is None: + return None + spans = first["spans"] # rebind-ok: accumulate spans across cursor pages + cursor = first.get("next_cursor") # rebind-ok: advance the trace cursor + while cursor: + page = await receiver.get_trace(trace_id, ALL_TRACES, trace_ref, cursor, LENS_DATASET_TRACE_PAGE_SIZE) + if page is None: + return None + spans = (*spans, *page["spans"]) + cursor = page.get("next_cursor") + return Trace(summary=first["summary"], agents=first["agents"], spans=spans, next_cursor=None) + + async def span(self, trace_id: str, span_id: str, trace_ref: str) -> SpanDetail | None: + try: + return await require_receiver(self.receiver).get_span(trace_id, span_id, ALL_TRACES, trace_ref) + except (TraceChanged, ValueError, OverflowError, RuntimeError) as error: + raise read_failure(error) from error + + async def findings(self, lens_id: str, ids: tuple[str, ...]) -> tuple[Finding, ...]: + lens: Final = await self.lenses.get(lens_id) + if lens is None or not can_access(self.scope, lens.scope): + raise HTTPException(404, "Lens not found") + wanted: Final = frozenset(ids) + return tuple(f for f in lens.findings if f.id in wanted) + + +async def get_dataset(datasets: DatasetStore, dataset_id: str, scope: Scope, revision: int | None = None) -> Dataset: + dataset: Final = await datasets.get(dataset_id, revision) + if dataset is None or not can_access(scope, Scope(team_id=dataset.team_id)): + raise HTTPException(404, "Dataset not found") + return dataset + + +@router.get("", response_model=tuple[DatasetSummary, ...]) +async def list_datasets(auth: Auth, datasets: Datasets) -> tuple[DatasetSummary, ...]: + scope: Final = user_scope(auth) + return tuple(s.summary for s in await datasets.summaries() if can_access(scope, Scope(team_id=s.team_id))) + + +@router.post("", response_model=Dataset) +async def create_dataset(body: DatasetCreate, auth: Auth, datasets: Datasets) -> Dataset: + user_scope(auth, write=True) + now: Final = datetime.now(timezone.utc) + dataset: Final = Dataset( + id=str(uuid4()), + name=body.name, + agent_name=body.agent_name, + team_id=auth.team_id or "", + created_at=now, + revision=0, + created_by=auth.user_id or "", + cases=(), + ) + if not await datasets.insert(dataset, now): + raise HTTPException(409, "Dataset already exists") + return dataset + + +@router.post("/build", response_model=BuildResult) +async def build_dataset_cases( + body: BuildRequest, auth: Auth, datasets: Datasets, lenses: Lenses, receiver: Receiver +) -> BuildResult: + scope: Final = user_scope(auth) + existing: Final[tuple[DatasetCase, ...]] = ( + (await get_dataset(datasets, body.dataset_id, scope)).cases if body.dataset_id else () + ) + return await build_cases(body, ProxyDatasetReader(receiver, lenses, scope), existing) + + +@router.get("/{dataset_id}", response_model=Dataset) +async def read_dataset(dataset_id: str, auth: Auth, datasets: Datasets, revision: RevisionQuery = None) -> Dataset: + return await get_dataset(datasets, dataset_id, user_scope(auth), revision) + + +@router.post("/{dataset_id}/revisions", response_model=Dataset) +async def save_revision(dataset_id: str, body: RevisionSave, auth: Auth, datasets: Datasets) -> Dataset: + latest: Final = await get_dataset(datasets, dataset_id, user_scope(auth, write=True)) + if body.base_revision != latest.revision: + raise HTTPException(409, "Dataset changed, reload") + cases: Final = revision_cases(body.cases) + if problem := revision_problem(cases): + raise HTTPException(422, problem) + saved: Final = latest.model_copy( + update=MappingProxyType({"revision": latest.revision + 1, "cases": cases, "created_by": auth.user_id or ""}) + ) + if not await datasets.insert(saved, datetime.now(timezone.utc)): + raise HTTPException(409, "Dataset changed, reload") + return saved + + +@router.get( + "/{dataset_id}/export", + response_class=Response, + responses={200: {"content": {"application/x-ndjson": {}}}}, +) +async def export_dataset(dataset_id: str, auth: Auth, datasets: Datasets, revision: RevisionQuery = None) -> Response: + dataset: Final = await get_dataset(datasets, dataset_id, user_scope(auth), revision) + return Response( + content=export_jsonl(dataset.cases), + media_type="application/x-ndjson", + headers=MappingProxyType( + {"Content-Disposition": f'attachment; filename="dataset-{dataset.id}-r{dataset.revision}.jsonl"'} + ), + ) + + +@router.get("/{dataset_id}/revisions/{revision}/cases", response_model=EvalCases) +async def eval_cases(dataset_id: str, revision: int, auth: Auth, datasets: Datasets) -> EvalCases: + dataset: Final = await get_dataset(datasets, dataset_id, user_scope(auth), revision) + return EvalCases(dataset_id=dataset.id, revision=dataset.revision, cases=included_cases(dataset.cases)) diff --git a/litellm/proxy/lens/dataset_repository.py b/litellm/proxy/lens/dataset_repository.py new file mode 100644 index 00000000000..50e5f3ee53a --- /dev/null +++ b/litellm/proxy/lens/dataset_repository.py @@ -0,0 +1,65 @@ +from datetime import datetime +from typing import Final, Protocol + +from pydantic import TypeAdapter + +from litellm.proxy.lens.models import Dataset, DatasetSummary, Record +from litellm.proxy.lens.repository import Database, Row + +_ROWS: Final = TypeAdapter(tuple[Row, ...]) + + +class StoredSummary(Record): + team_id: str + summary: DatasetSummary + + +class DatasetStore(Protocol): + async def summaries(self) -> tuple[StoredSummary, ...]: ... + async def get(self, dataset_id: str, revision: int | None = None) -> Dataset | None: ... + async def insert(self, dataset: Dataset, saved_at: datetime) -> bool: ... + + +class DatasetRepository: + def __init__(self, db: Database) -> None: + self.db: Final = db + + async def summaries(self) -> tuple[StoredSummary, ...]: + rows: Final = _ROWS.validate_python( + await self.db.query_raw( + """SELECT jsonb_build_object( + 'team_id', data->>'team_id', + 'summary', jsonb_build_object( + 'id', id, 'name', data->>'name', 'agent_name', data->>'agent_name', 'revision', revision, + 'case_count', jsonb_array_length(data->'cases'), + 'updated_at', created_at AT TIME ZONE 'UTC' + ) + ) AS data FROM ( + SELECT DISTINCT ON (id) id, revision, created_at, data FROM "LiteLLM_LensDataset" + ORDER BY id, revision DESC + ) AS latest ORDER BY created_at DESC""" + ) + ) + return tuple(StoredSummary.model_validate(row.data) for row in rows) + + async def get(self, dataset_id: str, revision: int | None = None) -> Dataset | None: + rows: Final = _ROWS.validate_python( + await self.db.query_raw( + """SELECT data FROM "LiteLLM_LensDataset" WHERE id=$1 AND ($2::int IS NULL OR revision=$2::int) + ORDER BY revision DESC LIMIT 1""", + dataset_id, + revision, + ) + ) + return Dataset.model_validate(rows[0].data) if rows else None + + async def insert(self, dataset: Dataset, saved_at: datetime) -> bool: + inserted: Final = await self.db.execute_raw( + """INSERT INTO "LiteLLM_LensDataset" (id, revision, created_at, data) + VALUES ($1, $2, $3::timestamptz AT TIME ZONE 'UTC', $4::jsonb) ON CONFLICT (id, revision) DO NOTHING""", + dataset.id, + dataset.revision, + saved_at.isoformat(), + dataset.model_dump_json(), + ) + return inserted == 1 diff --git a/litellm/proxy/lens/datasets.py b/litellm/proxy/lens/datasets.py new file mode 100644 index 00000000000..aef672833df --- /dev/null +++ b/litellm/proxy/lens/datasets.py @@ -0,0 +1,307 @@ +import hashlib +from collections.abc import Iterable +from dataclasses import dataclass +from functools import partial, reduce +from itertools import chain +from types import MappingProxyType +from typing import Final, Protocol, TypeAlias + +from pydantic import BaseModel, ConfigDict, ValidationError +from typing_extensions import assert_never + +from litellm.constants import LENS_DATASET_MAX_CASE_CHARS, LENS_DATASET_MAX_CASES +from litellm.proxy.lens.models import ( + BuildRequest, + BuildResult, + BuildSource, + CaseSource, + DatasetCase, + DatasetMessage, + DatasetToolCall, + Finding, + FindingSource, + Record, + SkippedCase, + TextSource, + TraceSource, +) +from litellm.proxy.lens.sources import parse_execution +from litellm.rust_bridge.trace.generated.types import SpanDetail, Trace, UIContent, UIMessage + +AGENT_VERSION_ATTRIBUTE: Final = "agent.version" + + +class DatasetReader(Protocol): + async def trace(self, trace_id: str, trace_ref: str) -> Trace | None: ... + async def span(self, trace_id: str, span_id: str, trace_ref: str) -> SpanDetail | None: ... + async def findings(self, lens_id: str, ids: tuple[str, ...]) -> tuple[Finding, ...]: ... + + +class _Content(Record): + messages: tuple[DatasetMessage, ...] + reply: str + tool_calls: tuple[DatasetToolCall, ...] + + +class _TextLine(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + messages: tuple[DatasetMessage, ...] + reply: str = "" + tool_calls: tuple[DatasetToolCall, ...] = () + expected: str = "" + source: CaseSource = CaseSource() + agent_version: str = "" + + +@dataclass(frozen=True, slots=True) +class _Reply: + text: str + tool_calls: tuple[DatasetToolCall, ...] + + +@dataclass(frozen=True, slots=True) +class _Admission: + seen: frozenset[str] + cases: tuple[DatasetCase, ...] = () + skipped: tuple[SkippedCase, ...] = () + + +Candidate: TypeAlias = DatasetCase | SkippedCase + + +def case_id(messages: tuple[DatasetMessage, ...], reply: str, tool_calls: tuple[DatasetToolCall, ...]) -> str: + content: Final = _Content(messages=messages, reply=reply, tool_calls=tool_calls) + return hashlib.sha256(content.model_dump_json().encode()).hexdigest() + + +def _tool_call_chars(tool_calls: tuple[DatasetToolCall, ...]) -> int: + return sum(len(t.name) + len(t.arguments) for t in tool_calls) + + +def case_chars(case: DatasetCase) -> int: + return ( + sum(len(m.content) + len(m.name) + _tool_call_chars(m.tool_calls) for m in case.messages) + + len(case.reply) + + _tool_call_chars(case.tool_calls) + + len(case.expected) + ) + + +def make_case( + messages: tuple[DatasetMessage, ...], + reply: str, + tool_calls: tuple[DatasetToolCall, ...], + source: CaseSource, + expected: str = "", + agent_version: str = "", +) -> Candidate: + if not messages and not reply.strip() and not tool_calls: + return SkippedCase(source=source, reason="no_content") + case: Final = DatasetCase( + id=case_id(messages, reply, tool_calls), + messages=messages, + reply=reply, + tool_calls=tool_calls, + expected=expected, + source=source, + agent_version=agent_version, + ) + if case_chars(case) > LENS_DATASET_MAX_CASE_CHARS: + return SkippedCase(source=source, reason="too_large") + return case + + +def _message(message: UIMessage) -> DatasetMessage: + return DatasetMessage( + role=message["role"], + content=message["content"], + name=message.get("name") or "", + tool_calls=_tool_calls((message,)), + ) + + +def _tool_calls(messages: Iterable[UIMessage]) -> tuple[DatasetToolCall, ...]: + calls: Final = chain.from_iterable(m.get("tool_calls", ()) for m in messages) + return tuple(DatasetToolCall(name=c["name"], arguments=c["arguments"]) for c in calls) + + +def _conversation(ui: UIContent, raw: str) -> tuple[DatasetMessage, ...]: + if ui["kind"] == "messages": + return tuple(_message(m) for m in ui["messages"]) + return (DatasetMessage(role="user", content=raw),) if raw.strip() else () + + +def _reply(ui: UIContent, raw: str) -> _Reply: + if ui["kind"] == "messages": + replies: Final = tuple(m for m in ui["messages"] if m["role"] == "assistant") + return _Reply("\n\n".join(m["content"] for m in replies if m["content"]), _tool_calls(replies)) + if ui["kind"] == "text": + return _Reply(ui["text"], ()) + return _Reply(raw, ()) + + +def case_from_span(detail: SpanDetail, source: CaseSource) -> Candidate: + reply: Final = _reply(detail["output_ui"], detail["output"]) + return make_case( + _conversation(detail["input_ui"], detail["input"]), + reply.text, + reply.tool_calls, + source.model_copy(update=MappingProxyType({"span_id": detail["span_id"]})), + agent_version=detail["attributes"].get(AGENT_VERSION_ATTRIBUTE, ""), + ) + + +async def _last_conversation(reader: DatasetReader, source: TraceSource, trace: Trace) -> SpanDetail | None: + llm_spans: Final = sorted( + (s for s in trace["spans"] if s["type"] == "llm"), key=lambda s: s["start_offset_ms"], reverse=True + ) + for span in llm_spans: + detail = await reader.span(source.trace_id, span["span_id"], source.trace_ref) + if detail is not None and detail["input_ui"]["kind"] == "messages": + return detail + return None + + +async def _trace_cases(reader: DatasetReader, source: TraceSource) -> tuple[Candidate, ...]: + origin: Final = CaseSource(trace_id=source.trace_id, trace_ref=source.trace_ref, span_id=source.span_id) + if source.span_id: + span: Final = await reader.span(source.trace_id, source.span_id, source.trace_ref) + return (case_from_span(span, origin) if span else SkippedCase(source=origin, reason="no_content"),) + trace: Final = await reader.trace(source.trace_id, source.trace_ref) + detail: Final = await _last_conversation(reader, source, trace) if trace else None + return (case_from_span(detail, origin) if detail else SkippedCase(source=origin, reason="no_content"),) + + +async def _evidence_case(reader: DatasetReader, origin: CaseSource, execution: str) -> Candidate: + try: + _, _, trace_id, trace_ref = parse_execution(execution) + except (ValueError, ValidationError): + return SkippedCase(source=origin, reason="no_content") + located: Final = origin.model_copy(update=MappingProxyType({"trace_id": trace_id, "trace_ref": trace_ref})) + span: Final = await reader.span(trace_id, origin.span_id, trace_ref) + return case_from_span(span, located) if span else SkippedCase(source=located, reason="no_content") + + +def _first_by_key(keys: tuple[str, ...]) -> frozenset[int]: + first: Final = MappingProxyType({key: index for index, key in reversed(tuple(enumerate(keys)))}) + return frozenset(first.values()) + + +def _evidence_spans(lens_id: str, findings: tuple[Finding, ...]) -> tuple[tuple[str, CaseSource], ...]: + pairs: Final = tuple(chain.from_iterable(((f.id, e) for e in f.evidence) for f in findings)) + kept: Final = _first_by_key(tuple(f"{e.execution_id}\0{e.span_id}" for _, e in pairs)) + return tuple( + (e.execution_id, CaseSource(span_id=e.span_id, finding_id=fid, lens_id=lens_id)) + for index, (fid, e) in enumerate(pairs) + if index in kept + ) + + +async def _finding_cases(reader: DatasetReader, source: FindingSource) -> tuple[Candidate, ...]: + findings: Final = await reader.findings(source.lens_id, source.finding_ids) + spans: Final = _evidence_spans(source.lens_id, findings) + return tuple([await _evidence_case(reader, origin, execution) for execution, origin in spans]) + + +def _looks_like_json_line(line: str) -> bool: + return line.lstrip().startswith("{") + + +def _text_line_case(line: str) -> Candidate: + try: + parsed: Final = _TextLine.model_validate_json(line) + except ValidationError: + return SkippedCase(source=CaseSource(), reason="invalid") + return make_case( + parsed.messages, + parsed.reply, + parsed.tool_calls, + parsed.source, + expected=parsed.expected, + agent_version=parsed.agent_version, + ) + + +def _text_cases(source: TextSource) -> tuple[Candidate, ...]: + lines: Final = tuple(line for line in source.text.splitlines() if line.strip()) + if lines and all(_looks_like_json_line(line) for line in lines): + return tuple(_text_line_case(line) for line in lines) + return (make_case((DatasetMessage(role="user", content=source.text),), "", (), CaseSource()),) + + +async def _source_cases(reader: DatasetReader, source: BuildSource) -> tuple[Candidate, ...]: + match source: + case TraceSource(): + return await _trace_cases(reader, source) + case FindingSource(): + return await _finding_cases(reader, source) + case TextSource(): + return _text_cases(source) + return assert_never(source) + + +def _skip(state: _Admission, skipped: SkippedCase) -> _Admission: + return _Admission(state.seen, state.cases, (*state.skipped, skipped)) + + +def _admit(existing_count: int, state: _Admission, candidate: Candidate) -> _Admission: + if isinstance(candidate, SkippedCase): + return _skip(state, candidate) + if candidate.id in state.seen: + return _skip(state, SkippedCase(source=candidate.source, reason="duplicate")) + if existing_count + len(state.cases) >= LENS_DATASET_MAX_CASES: + return _skip(state, SkippedCase(source=candidate.source, reason="over_limit")) + return _Admission(state.seen | {candidate.id}, (*state.cases, candidate), state.skipped) + + +def _unread_source(source: BuildSource) -> CaseSource: + match source: + case TraceSource(): + return CaseSource(trace_id=source.trace_id, trace_ref=source.trace_ref, span_id=source.span_id) + case FindingSource(): + return CaseSource(lens_id=source.lens_id, finding_id=source.finding_ids[0]) + case TextSource(): + return CaseSource() + return assert_never(source) + + +async def _admit_source( + reader: DatasetReader, existing_count: int, state: _Admission, source: BuildSource +) -> _Admission: + if existing_count + len(state.cases) >= LENS_DATASET_MAX_CASES: + return _skip(state, SkippedCase(source=_unread_source(source), reason="over_limit")) + return reduce(partial(_admit, existing_count), await _source_cases(reader, source), state) + + +async def build_cases(request: BuildRequest, reader: DatasetReader, existing: tuple[DatasetCase, ...]) -> BuildResult: + state = _Admission(seen=frozenset(c.id for c in existing)) # rebind-ok: sources are read in order until full + for source in request.sources: + state = await _admit_source(reader, len(existing), state, source) + return BuildResult(cases=state.cases, skipped=state.skipped) + + +def rehashed(case: DatasetCase) -> DatasetCase: + return case.model_copy(update=MappingProxyType({"id": case_id(case.messages, case.reply, case.tool_calls)})) + + +def revision_cases(cases: tuple[DatasetCase, ...]) -> tuple[DatasetCase, ...]: + hashed: Final = tuple(rehashed(c) for c in cases) + kept: Final = _first_by_key(tuple(c.id for c in hashed)) + return tuple(c for index, c in enumerate(hashed) if index in kept) + + +def revision_problem(cases: tuple[DatasetCase, ...]) -> str | None: + if len(cases) > LENS_DATASET_MAX_CASES: + return f"A dataset holds at most {LENS_DATASET_MAX_CASES} cases" + if any(case_chars(c) > LENS_DATASET_MAX_CASE_CHARS for c in cases): + return f"Each case must be at most {LENS_DATASET_MAX_CASE_CHARS} characters" + return None + + +def included_cases(cases: tuple[DatasetCase, ...]) -> tuple[DatasetCase, ...]: + return tuple(c for c in cases if c.included) + + +def export_jsonl(cases: tuple[DatasetCase, ...]) -> str: + return "".join(c.model_dump_json() + "\n" for c in included_cases(cases)) diff --git a/litellm/proxy/lens/models.py b/litellm/proxy/lens/models.py index b29809a440b..7d3f791e55a 100644 --- a/litellm/proxy/lens/models.py +++ b/litellm/proxy/lens/models.py @@ -500,3 +500,106 @@ class ModelResult(Record): cost: float context_exceeded: bool = False finish_reason: Literal["length", "content_filter"] | None = Field(default=None, exclude=True) + + +class DatasetToolCall(Record): + name: str + arguments: str + + +class DatasetMessage(Record): + role: Literal["system", "user", "assistant", "tool"] + content: str + name: str = "" + tool_calls: tuple[DatasetToolCall, ...] = () + + +class CaseSource(Record): + trace_id: str = "" + trace_ref: str = "" + span_id: str = "" + finding_id: str = "" + lens_id: str = "" + + +class DatasetCase(Record): + id: str + messages: tuple[DatasetMessage, ...] + reply: str = "" + tool_calls: tuple[DatasetToolCall, ...] = () + expected: str = "" + included: bool = True + source: CaseSource + agent_version: str = "" + + +class SkippedCase(Record): + source: CaseSource + reason: Literal["duplicate", "no_content", "too_large", "over_limit", "invalid"] + + +class TraceSource(Record): + kind: Literal["trace"] = "trace" + trace_id: str = Field(min_length=1) + trace_ref: str = "" + span_id: str = "" + + +class FindingSource(Record): + kind: Literal["finding"] = "finding" + lens_id: str = Field(min_length=1) + finding_ids: tuple[str, ...] = Field(min_length=1) + + +class TextSource(Record): + kind: Literal["text"] = "text" + text: str = Field(min_length=1) + + +BuildSource: TypeAlias = Annotated[TraceSource | FindingSource | TextSource, Field(discriminator="kind")] + + +class BuildRequest(Record): + sources: tuple[BuildSource, ...] = Field(min_length=1) + dataset_id: str = "" + + +class BuildResult(Record): + cases: tuple[DatasetCase, ...] + skipped: tuple[SkippedCase, ...] + + +class DatasetCreate(Record): + name: str = Field(min_length=1, max_length=120) + agent_name: str = "" + + +class Dataset(Record): + id: str + name: str + agent_name: str + team_id: str + created_at: datetime + revision: int + created_by: str + cases: tuple[DatasetCase, ...] + + +class DatasetSummary(Record): + id: str + name: str + agent_name: str + revision: int + case_count: int + updated_at: datetime + + +class RevisionSave(Record): + base_revision: int = Field(ge=0) + cases: tuple[DatasetCase, ...] + + +class EvalCases(Record): + dataset_id: str + revision: int + cases: tuple[DatasetCase, ...] diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5d6ba0f8e13..1f3f6ea0adb 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -579,6 +579,7 @@ from litellm.proxy.hooks.prompt_injection_detection import ( ) from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_spend_event 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.list_api.common import ( ManagementProblem, @@ -20110,6 +20111,7 @@ app.include_router(auto_router_management_router) app.include_router(tag_management_router) app.include_router(workflow_management_router) app.include_router(memory_router) +app.include_router(lens_dataset_router) app.include_router(lens_router) app.include_router(plugin_router) app.include_router(cost_tracking_settings_router) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 4aabd56131b..514df905866 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1965,3 +1965,12 @@ model LiteLLM_LensWorker { token_hash String @unique data Json } + +model LiteLLM_LensDataset { + id String + revision Int + created_at DateTime + data Json + + @@id([id, revision]) +} diff --git a/schema.prisma b/schema.prisma index 4aabd56131b..514df905866 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1965,3 +1965,12 @@ model LiteLLM_LensWorker { token_hash String @unique data Json } + +model LiteLLM_LensDataset { + id String + revision Int + created_at DateTime + data Json + + @@id([id, revision]) +} diff --git a/tests/integration/database/test_lens_dataset_repository.py b/tests/integration/database/test_lens_dataset_repository.py new file mode 100644 index 00000000000..baeaa4d50b4 --- /dev/null +++ b/tests/integration/database/test_lens_dataset_repository.py @@ -0,0 +1,149 @@ +import os +from collections.abc import AsyncIterator +from datetime import datetime, timedelta, timezone +from typing import Final +from uuid import uuid4 + +import pytest +import pytest_asyncio +from prisma import Prisma +from prisma.types import DatasourceOverride + +from litellm.proxy.db.prisma_client import PrismaWrapper +from litellm.proxy.lens.dataset_repository import DatasetRepository, StoredSummary +from litellm.proxy.lens.models import CaseSource, Dataset, DatasetCase, DatasetMessage, DatasetSummary +from litellm.proxy.lens.repository import WriterDatabase + +SAVED_AT: Final = datetime(2026, 3, 1, 12, 0, 0, 123000, tzinfo=timezone.utc) + + +@pytest_asyncio.fixture(loop_scope="function") +async def lens_db() -> AsyncIterator[Prisma]: + async with Prisma(datasource=DatasourceOverride(url=os.environ["DATABASE_URL"])) as db: + yield db + + +@pytest_asyncio.fixture(loop_scope="function") +async def dataset_ids(lens_db: Prisma) -> AsyncIterator[tuple[str, ...]]: + ids: Final = tuple(uuid4().hex for _ in range(3)) + yield ids + await lens_db.execute_raw('DELETE FROM "LiteLLM_LensDataset" WHERE id = ANY($1::text[])', list(ids)) + + +def _repo(db: Prisma) -> DatasetRepository: + return DatasetRepository(WriterDatabase(PrismaWrapper(db))) + + +def _dataset(dataset_id: str, revision: int, case_count: int, team_id: str = "team-a") -> Dataset: + return Dataset( + id=dataset_id, + name=f"Dataset r{revision}", + agent_name="support-agent", + team_id=team_id, + created_at=SAVED_AT, + revision=revision, + created_by="user-1", + cases=tuple( + DatasetCase( + id=f"case-{revision}-{i}", + messages=(DatasetMessage(role="user", content=f"question {i}"),), + reply=f"answer {i}", + source=CaseSource(trace_id=f"trace-{i}"), + ) + for i in range(case_count) + ), + ) + + +@pytest.mark.asyncio +async def test_get_returns_requested_revision_and_defaults_to_latest( + lens_db: Prisma, dataset_ids: tuple[str, ...] +) -> None: + repo: Final = _repo(lens_db) + dataset_id: Final = dataset_ids[0] + revisions: Final = tuple(_dataset(dataset_id, revision, revision) for revision in (1, 3, 2)) + assert [await repo.insert(dataset, SAVED_AT) for dataset in revisions] == [True, True, True] + + assert await repo.get(dataset_id, revision=1) == revisions[0] + assert await repo.get(dataset_id, revision=2) == revisions[2] + assert await repo.get(dataset_id) == revisions[1] + + +@pytest.mark.asyncio +async def test_get_unknown_dataset_or_revision_returns_none(lens_db: Prisma, dataset_ids: tuple[str, ...]) -> None: + repo: Final = _repo(lens_db) + assert await repo.insert(_dataset(dataset_ids[0], 1, 1), SAVED_AT) + + assert await repo.get(dataset_ids[1]) is None + assert await repo.get(dataset_ids[1], revision=1) is None + assert await repo.get(dataset_ids[0], revision=2) is None + + +@pytest.mark.asyncio +async def test_inserting_an_existing_revision_is_rejected_and_keeps_the_first( + lens_db: Prisma, dataset_ids: tuple[str, ...] +) -> None: + repo: Final = _repo(lens_db) + original: Final = _dataset(dataset_ids[0], 1, 1) + overwrite: Final = _dataset(dataset_ids[0], 1, 4).model_copy(update={"name": "Overwritten"}) + + assert await repo.insert(original, SAVED_AT) is True + assert await repo.insert(overwrite, SAVED_AT + timedelta(days=1)) is False + + assert await repo.get(dataset_ids[0], revision=1) == original + summaries: Final = tuple(s for s in await repo.summaries() if s.summary.id == dataset_ids[0]) + assert tuple(s.summary.updated_at for s in summaries) == (SAVED_AT,) + + +@pytest.mark.asyncio +async def test_summaries_list_each_dataset_once_at_its_latest_revision_newest_first( + lens_db: Prisma, dataset_ids: tuple[str, ...] +) -> None: + repo: Final = _repo(lens_db) + older, newer, single = dataset_ids + writes: Final = ( + (_dataset(older, 1, 1, "team-a"), SAVED_AT), + (_dataset(older, 2, 3, "team-a"), SAVED_AT + timedelta(minutes=1)), + (_dataset(newer, 1, 5, "team-b"), SAVED_AT + timedelta(minutes=2)), + (_dataset(newer, 2, 2, "team-b"), SAVED_AT + timedelta(minutes=4)), + (_dataset(single, 1, 0, ""), SAVED_AT + timedelta(minutes=3)), + ) + assert [await repo.insert(dataset, saved_at) for dataset, saved_at in writes] == [True] * len(writes) + + summaries: Final = tuple(s for s in await repo.summaries() if s.summary.id in dataset_ids) + + assert summaries == ( + StoredSummary( + team_id="team-b", + summary=DatasetSummary( + id=newer, + name="Dataset r2", + agent_name="support-agent", + revision=2, + case_count=2, + updated_at=SAVED_AT + timedelta(minutes=4), + ), + ), + StoredSummary( + team_id="", + summary=DatasetSummary( + id=single, + name="Dataset r1", + agent_name="support-agent", + revision=1, + case_count=0, + updated_at=SAVED_AT + timedelta(minutes=3), + ), + ), + StoredSummary( + team_id="team-a", + summary=DatasetSummary( + id=older, + name="Dataset r2", + agent_name="support-agent", + revision=2, + case_count=3, + updated_at=SAVED_AT + timedelta(minutes=1), + ), + ), + ) diff --git a/tests/unit/proxy/lens/test_dataset_endpoints.py b/tests/unit/proxy/lens/test_dataset_endpoints.py new file mode 100644 index 00000000000..53d60d98d20 --- /dev/null +++ b/tests/unit/proxy/lens/test_dataset_endpoints.py @@ -0,0 +1,408 @@ +import json +from collections.abc import Mapping +from datetime import datetime +from typing import Final + +import pytest +from fastapi import HTTPException + +from litellm.constants import LENS_DATASET_MAX_CASE_CHARS, LENS_DATASET_MAX_CASES +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.lens.dataset_endpoints import ( + ProxyDatasetReader, + build_dataset_cases, + create_dataset, + eval_cases, + export_dataset, + list_datasets, + read_dataset, + save_revision, +) +from litellm.proxy.lens.dataset_repository import StoredSummary +from litellm.proxy.lens.datasets import case_id +from litellm.proxy.lens.models import ( + BuildRequest, + CaseSource, + Dataset, + DatasetCase, + DatasetCreate, + DatasetMessage, + DatasetSummary, + Evidence, + Lens, + RevisionSave, + Scope, + TextSource, +) +from litellm.proxy.lens.repository import LensRepository, Row +from litellm.rust_bridge.trace.errors import TraceChanged +from litellm.rust_bridge.trace.generated.types import SpanDetail, Trace, TraceScope +from litellm.rust_bridge.trace.storage import ClickHouseStorage +from litellm.tracing import TraceReceiver +from tests.unit.proxy.lens.test_datasets import detail, span, stored_finding, trace +from tests.unit.proxy.lens.test_state import lens + +ADMIN: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin", team_id="alpha") + + +class MemoryStore: + def __init__(self) -> None: + self.rows: Final[dict[tuple[str, int], Dataset]] = {} # mutable-ok: stands in for the table + + async def summaries(self) -> tuple[StoredSummary, ...]: + latest: Final = {i: d for (i, _), d in sorted(self.rows.items(), key=lambda item: item[0][1])} + return tuple(StoredSummary(team_id=d.team_id, summary=summary(d)) for d in latest.values()) + + async def get(self, dataset_id: str, revision: int | None = None) -> Dataset | None: + revisions: Final = sorted(r for i, r in self.rows if i == dataset_id) + wanted: Final = revision if revision is not None else (revisions[-1] if revisions else None) + return self.rows.get((dataset_id, wanted)) if wanted is not None else None + + async def insert(self, dataset: Dataset, saved_at: datetime) -> bool: + key: Final = (dataset.id, dataset.revision) + if key in self.rows: + return False + self.rows[key] = dataset + return True + + +def case(text: str, included: bool = True) -> DatasetCase: + message: Final = (DatasetMessage(role="user", content=text),) + return DatasetCase(id=case_id(message, "", ()), messages=message, included=included, source=CaseSource()) + + +@pytest.mark.asyncio +async def test_saving_inserts_a_new_revision_and_leaves_the_older_one_unchanged() -> None: + store: Final = MemoryStore() + created: Final = await create_dataset_named(store) + first: Final = await save_revision(created.id, RevisionSave(base_revision=0, cases=(case("a"),)), ADMIN, store) + second: Final = await save_revision( + created.id, RevisionSave(base_revision=1, cases=(case("a"), case("b"))), ADMIN, store + ) + + assert (first.revision, second.revision) == (1, 2) + assert await read_dataset(created.id, ADMIN, store, revision=1) == first + assert (await read_dataset(created.id, ADMIN, store, revision=0)).cases == () + assert await read_dataset(created.id, ADMIN, store) == second + + +@pytest.mark.asyncio +async def test_saving_on_a_stale_revision_is_a_conflict_and_writes_nothing() -> None: + store: Final = MemoryStore() + created: Final = await create_dataset_named(store) + await save_revision(created.id, RevisionSave(base_revision=0, cases=(case("a"),)), ADMIN, store) + + with pytest.raises(HTTPException) as error: + await save_revision(created.id, RevisionSave(base_revision=0, cases=(case("b"),)), ADMIN, store) + assert error.value.status_code == 409 + assert sorted(store.rows) == [(created.id, 0), (created.id, 1)] + + +@pytest.mark.asyncio +async def test_saving_dedupes_cases_and_restores_content_hash_ids_but_keeps_edits() -> None: + store: Final = MemoryStore() + created: Final = await create_dataset_named(store) + edited: Final = case("a").model_copy(update={"id": "forged", "expected": "say hi"}) + saved: Final = await save_revision( + created.id, RevisionSave(base_revision=0, cases=(edited, case("a"))), ADMIN, store + ) + + assert len(saved.cases) == 1 + assert saved.cases[0].id == case("a").id + assert saved.cases[0].expected == "say hi" + + +@pytest.mark.asyncio +async def test_eval_cases_return_only_included_cases_of_the_requested_revision() -> None: + store: Final = MemoryStore() + created: Final = await create_dataset_named(store) + await save_revision( + created.id, RevisionSave(base_revision=0, cases=(case("a"), case("b", included=False))), ADMIN, store + ) + await save_revision(created.id, RevisionSave(base_revision=1, cases=(case("c"),)), ADMIN, store) + + cases: Final = await eval_cases(created.id, 1, ADMIN, store) + assert (cases.dataset_id, cases.revision) == (created.id, 1) + assert tuple(c.messages[0].content for c in cases.cases) == ("a",) + + +@pytest.mark.asyncio +async def test_read_only_admin_can_read_but_cannot_save_a_revision() -> None: + store: Final = MemoryStore() + created: Final = await create_dataset_named(store) + viewer: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + + assert (await read_dataset(created.id, viewer, store)).id == created.id + with pytest.raises(HTTPException) as error: + await save_revision(created.id, RevisionSave(base_revision=0, cases=(case("a"),)), viewer, store) + assert error.value.status_code == 403 + + +async def create_dataset_named(store: MemoryStore) -> Dataset: + return await create_dataset(DatasetCreate(name="Refunds", agent_name="support"), ADMIN, store) + + +def summary(dataset: Dataset) -> DatasetSummary: + return DatasetSummary( + id=dataset.id, + name=dataset.name, + agent_name=dataset.agent_name, + revision=dataset.revision, + case_count=len(dataset.cases), + updated_at=dataset.created_at, + ) + + +class RejectingStore: + def __init__(self, stored: Dataset | None = None) -> None: + self.stored: Final = stored + + async def summaries(self) -> tuple[StoredSummary, ...]: + return () + + async def get(self, dataset_id: str, revision: int | None = None) -> Dataset | None: + return self.stored + + async def insert(self, dataset: Dataset, saved_at: datetime) -> bool: + return False + + +class TraceStorage(ClickHouseStorage): + def __init__( + self, + pages: Mapping[str | None, Trace], + spans: Mapping[str, SpanDetail] | None = None, + error: Exception | None = None, + ) -> None: + self.pages: Final = pages + self.spans: Final = spans or {} + self.error: Final = error + + async def get_trace( + self, + trace_id: str, + scope: TraceScope, + trace_ref: str = "", + cursor: str | None = None, + page_size: int | None = None, + ) -> Trace | None: + if self.error: + raise self.error + return self.pages.get(cursor) + + async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None: + if self.error: + raise self.error + return self.spans.get(span_id) + + +class LensTable: + def __init__(self, stored: Lens) -> None: + self.stored: Final = stored + + async def query_raw(self, query: str, *args: object) -> tuple[Row, ...]: + return (Row(data=self.stored.model_dump(mode="json")),) if args == (self.stored.id,) else () + + async def execute_raw(self, query: str, *args: object) -> int: + return 0 + + +def page(next_cursor: str | None, *span_ids: str) -> Trace: + return {**trace(*(span(s, "llm", i) for i, s in enumerate(span_ids))), "next_cursor": next_cursor} + + +def lenses(stored: Lens | None = None) -> LensRepository: + return LensRepository(LensTable(stored or lens())) + + +def reader( + storage: TraceStorage | None = None, stored: Lens | None = None, scope: Scope | None = None +) -> ProxyDatasetReader: + return ProxyDatasetReader( + TraceReceiver(storage) if storage else None, + lenses(stored), + scope or Scope(all_teams=True), + ) + + +@pytest.mark.asyncio +async def test_reading_a_trace_follows_every_cursor_page_and_joins_the_spans() -> None: + storage: Final = TraceStorage({None: page("c1", "a", "b"), "c1": page("c2", "c"), "c2": page(None, "d")}) + read: Final = await reader(storage).trace("t1", "ref") + + assert read is not None + assert tuple(s["span_id"] for s in read["spans"]) == ("a", "b", "c", "d") + assert read.get("next_cursor") is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("pages", ({}, {None: page("c1", "a")}), ids=("missing-trace", "page-disappears")) +async def test_reading_a_trace_is_none_when_the_trace_or_a_later_page_is_gone( + pages: Mapping[str | None, Trace], +) -> None: + assert await reader(TraceStorage(pages)).trace("t1", "ref") is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("error", "status"), ((TraceChanged("moved"), 409), (ValueError("bad cursor"), 400), (RuntimeError("down"), 503)) +) +async def test_trace_and_span_read_failures_become_the_matching_http_error(error: Exception, status: int) -> None: + failing: Final = reader(TraceStorage({}, error=error)) + + with pytest.raises(HTTPException) as traced: + await failing.trace("t1", "ref") + with pytest.raises(HTTPException) as spanned: + await failing.span("t1", "s1", "ref") + assert (traced.value.status_code, spanned.value.status_code) == (status, status) + + +@pytest.mark.asyncio +async def test_reading_traces_without_tracing_enabled_is_not_implemented() -> None: + disabled: Final = reader() + + with pytest.raises(HTTPException) as traced: + await disabled.trace("t1", "ref") + with pytest.raises(HTTPException) as spanned: + await disabled.span("t1", "s1", "ref") + assert (traced.value.status_code, spanned.value.status_code) == (501, 501) + + +@pytest.mark.asyncio +async def test_reading_a_span_returns_its_detail() -> None: + stored: Final = detail("s1", "refund?") + assert await reader(TraceStorage({}, {"s1": stored})).span("t1", "s1", "") == stored + + +def lens_with_findings() -> Lens: + evidence: Final = Evidence(execution_id="run", span_id="s1", quote="q") + return lens().model_copy( + update={ + "findings": (stored_finding("f1", evidence), stored_finding("f2", evidence), stored_finding("f3", evidence)) + } + ) + + +@pytest.mark.asyncio +async def test_findings_returns_only_the_requested_ids() -> None: + found: Final = await reader(stored=lens_with_findings()).findings("lens", ("f3", "f1", "unknown")) + assert tuple(f.id for f in found) == ("f1", "f3") + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("lens_id", "scope"), + (("missing", Scope(all_teams=True)), ("lens", Scope(team_id="beta"))), + ids=("unknown-lens", "lens-outside-scope"), +) +async def test_findings_of_an_unknown_or_inaccessible_lens_are_not_found(lens_id: str, scope: Scope) -> None: + with pytest.raises(HTTPException) as error: + await reader(stored=lens_with_findings(), scope=scope).findings(lens_id, ("f1",)) + assert error.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_list_returns_a_summary_of_the_latest_revision_of_each_dataset() -> None: + store: Final = MemoryStore() + first: Final = await create_dataset_named(store) + second: Final = await create_dataset_named(store) + await save_revision(first.id, RevisionSave(base_revision=0, cases=(case("a"), case("b"))), ADMIN, store) + + listed: Final = await list_datasets(ADMIN, store) + assert sorted((s.id, s.revision, s.case_count) for s in listed) == sorted(((first.id, 1, 2), (second.id, 0, 0))) + + +@pytest.mark.asyncio +async def test_listing_requires_proxy_admin_access() -> None: + with pytest.raises(HTTPException) as error: + await list_datasets(UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), MemoryStore()) + assert error.value.status_code == 403 + + +def text_source(*texts: str) -> TextSource: + return TextSource(text="\n".join(json.dumps({"messages": [{"role": "user", "content": t}]}) for t in texts)) + + +@pytest.mark.asyncio +async def test_building_into_a_dataset_skips_cases_it_already_holds() -> None: + store: Final = MemoryStore() + created: Final = await create_dataset_named(store) + await save_revision(created.id, RevisionSave(base_revision=0, cases=(case("a"),)), ADMIN, store) + request: Final = BuildRequest(sources=(text_source("a", "b"),), dataset_id=created.id) + + built: Final = await build_dataset_cases(request, ADMIN, store, lenses(), None) + fresh: Final = await build_dataset_cases( + request.model_copy(update={"dataset_id": ""}), ADMIN, store, lenses(), None + ) + + assert tuple(c.messages[0].content for c in built.cases) == ("b",) + assert tuple(s.reason for s in built.skipped) == ("duplicate",) + assert tuple(c.messages[0].content for c in fresh.cases) == ("a", "b") + + +@pytest.mark.asyncio +async def test_building_into_an_unknown_dataset_is_not_found() -> None: + request: Final = BuildRequest(sources=(text_source("a"),), dataset_id="missing") + with pytest.raises(HTTPException) as error: + await build_dataset_cases(request, ADMIN, MemoryStore(), lenses(), None) + assert error.value.status_code == 404 + + +def exported_texts(body: bytes) -> tuple[str, ...]: + return tuple(DatasetCase.model_validate_json(line).messages[0].content for line in body.splitlines()) + + +@pytest.mark.asyncio +async def test_export_downloads_included_cases_of_the_requested_revision_as_ndjson() -> None: + store: Final = MemoryStore() + created: Final = await create_dataset_named(store) + await save_revision( + created.id, RevisionSave(base_revision=0, cases=(case("a"), case("b", included=False))), ADMIN, store + ) + await save_revision(created.id, RevisionSave(base_revision=1, cases=(case("c"),)), ADMIN, store) + + exported: Final = await export_dataset(created.id, ADMIN, store, revision=1) + latest: Final = await export_dataset(created.id, ADMIN, store) + + assert exported.media_type == "application/x-ndjson" + assert exported.headers["content-disposition"] == f'attachment; filename="dataset-{created.id}-r1.jsonl"' + assert exported_texts(bytes(exported.body)) == ("a",) + assert latest.headers["content-disposition"] == f'attachment; filename="dataset-{created.id}-r2.jsonl"' + assert exported_texts(bytes(latest.body)) == ("c",) + + +@pytest.mark.asyncio +async def test_creating_a_dataset_whose_insert_is_rejected_is_a_conflict() -> None: + with pytest.raises(HTTPException) as error: + await create_dataset(DatasetCreate(name="Refunds"), ADMIN, RejectingStore()) + assert error.value.status_code == 409 + + +@pytest.mark.asyncio +async def test_saving_loses_to_a_concurrent_save_of_the_same_revision() -> None: + created: Final = await create_dataset_named(MemoryStore()) + + with pytest.raises(HTTPException) as error: + await save_revision( + created.id, RevisionSave(base_revision=0, cases=(case("a"),)), ADMIN, RejectingStore(created) + ) + assert error.value.status_code == 409 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "cases", + ( + tuple(case(str(i)) for i in range(LENS_DATASET_MAX_CASES + 1)), + (case("x" * (LENS_DATASET_MAX_CASE_CHARS + 1)),), + ), + ids=("too-many-cases", "case-too-large"), +) +async def test_saving_an_oversized_revision_is_rejected_and_writes_nothing(cases: tuple[DatasetCase, ...]) -> None: + store: Final = MemoryStore() + created: Final = await create_dataset_named(store) + + with pytest.raises(HTTPException) as error: + await save_revision(created.id, RevisionSave(base_revision=0, cases=cases), ADMIN, store) + assert error.value.status_code == 422 + assert sorted(store.rows) == [(created.id, 0)] diff --git a/tests/unit/proxy/lens/test_datasets.py b/tests/unit/proxy/lens/test_datasets.py new file mode 100644 index 00000000000..eda9b26927e --- /dev/null +++ b/tests/unit/proxy/lens/test_datasets.py @@ -0,0 +1,465 @@ +import json +from collections.abc import Mapping +from datetime import datetime, timezone +from typing import Final + +import pytest + +from litellm.constants import LENS_DATASET_MAX_CASE_CHARS, LENS_DATASET_MAX_CASES +from litellm.proxy.lens.datasets import build_cases, case_id, export_jsonl, revision_problem +from litellm.proxy.lens.models import ( + BuildRequest, + BuildResult, + CaseSource, + DatasetCase, + DatasetMessage, + DatasetToolCall, + Evidence, + Finding, + FindingSource, + SkippedCase, + TextSource, + TraceSource, +) +from litellm.proxy.lens.sources import execution_id +from litellm.rust_bridge.trace.generated.types import ( + Span, + SpanDetail, + SpanType, + Trace, + TraceSummary, + UIContent, + UIField, + UIFields, + UIMessage, + UIMessages, + UIText, +) + +NOW: Final = datetime(2026, 1, 15, tzinfo=timezone.utc) + + +def messages(*items: UIMessage) -> UIMessages: + return UIMessages(kind="messages", messages=items) + + +def detail( + span_id: str, + question: str, + answer: str = "Done", + attributes: Mapping[str, str] | None = None, + input_ui: UIContent | None = None, +) -> SpanDetail: + return SpanDetail( + span_id=span_id, + input_ui=input_ui + or messages(UIMessage(role="system", content="Be terse"), UIMessage(role="user", content=question)), + output_ui=messages( + UIMessage( + role="assistant", + content=answer, + tool_calls=({"name": "search", "arguments": json.dumps({"q": question})},), + ) + ), + input=question, + output=answer, + attributes=attributes or {}, + ) + + +def span(span_id: str, kind: SpanType, offset: float) -> Span: + return Span( + span_id=span_id, + parent_span_id=None, + name=span_id, + type=kind, + agent="support", + framework="", + start_offset_ms=offset, + duration_ms=1, + status="ok", + error=None, + error_truncated=False, + input_preview="", + model=None, + input_tokens=0, + output_tokens=0, + litellm_request_id=None, + spend=None, + ) + + +def trace(*spans: Span) -> Trace: + summary: Final = TraceSummary( + trace_id="t1", + name="run", + service="svc", + input_preview="", + start_time="", + duration_ms=1, + status="ok", + span_count=len(spans), + agent_count=1, + agent_invocations=1, + llm_calls=len(spans), + tool_calls=0, + error_count=0, + input_tokens=0, + output_tokens=0, + models=(), + spend=None, + ) + return Trace(summary=summary, agents=(), spans=spans) + + +def stored_finding(finding_id: str, *evidence: Evidence) -> Finding: + return Finding( + id=finding_id, + title="Repeated failed searches", + description="The agent repeats the same failed search", + check_id="retries", + evidence=evidence, + first_seen=NOW, + last_seen=NOW, + revision=1, + ) + + +class FakeReader: + def __init__( + self, + spans: Mapping[tuple[str, str], SpanDetail], + traces: Mapping[str, Trace] | None = None, + findings: tuple[Finding, ...] = (), + ) -> None: + self.spans: Final = spans + self.traces: Final = traces or {} + self.stored: Final = findings + + async def trace(self, trace_id: str, trace_ref: str) -> Trace | None: + return self.traces.get(trace_id) + + async def span(self, trace_id: str, span_id: str, trace_ref: str) -> SpanDetail | None: + return self.spans.get((trace_id, span_id)) + + async def findings(self, lens_id: str, ids: tuple[str, ...]) -> tuple[Finding, ...]: + return tuple(f for f in self.stored if f.id in ids) + + +async def build(reader: FakeReader, *sources: TraceSource | FindingSource | TextSource) -> BuildResult: + return await build_cases(BuildRequest(sources=sources), reader, ()) + + +@pytest.mark.asyncio +async def test_a_span_becomes_a_case_with_its_conversation_reply_and_tool_calls() -> None: + reader: Final = FakeReader({("t1", "s1"): detail("s1", "refund?", "Refunded", {"agent.version": "v7"})}) + result: Final = await build(reader, TraceSource(trace_id="t1", trace_ref="ref", span_id="s1")) + + case: Final = result.cases[0] + assert case.messages == ( + DatasetMessage(role="system", content="Be terse"), + DatasetMessage(role="user", content="refund?"), + ) + assert case.reply == "Refunded" + assert case.tool_calls == (DatasetToolCall(name="search", arguments=json.dumps({"q": "refund?"})),) + assert case.source == CaseSource(trace_id="t1", trace_ref="ref", span_id="s1") + assert case.agent_version == "v7" + assert case.id == case_id(case.messages, case.reply, case.tool_calls) + + +def history(tool: str) -> UIMessages: + return messages( + UIMessage(role="user", content="refund A1"), + UIMessage(role="assistant", content="", tool_calls=({"name": tool, "arguments": '{"id":"A1"}'},)), + UIMessage(role="tool", content="found"), + ) + + +@pytest.mark.asyncio +async def test_tool_calls_in_the_conversation_history_are_kept_and_change_the_case_id() -> None: + reader: Final = FakeReader( + { + ("t1", "s1"): detail("s1", "refund A1", input_ui=history("lookup_order")), + ("t1", "s2"): detail("s2", "refund A1", input_ui=history("cancel_order")), + } + ) + result: Final = await build( + reader, TraceSource(trace_id="t1", span_id="s1"), TraceSource(trace_id="t1", span_id="s2") + ) + + assert result.cases[0].messages[1].tool_calls == (DatasetToolCall(name="lookup_order", arguments='{"id":"A1"}'),) + assert result.cases[1].messages[1].tool_calls == (DatasetToolCall(name="cancel_order", arguments='{"id":"A1"}'),) + assert result.skipped == () + + +@pytest.mark.asyncio +async def test_agent_version_is_empty_when_the_span_does_not_report_one() -> None: + reader: Final = FakeReader({("t1", "s1"): detail("s1", "refund?")}) + result: Final = await build(reader, TraceSource(trace_id="t1", span_id="s1")) + assert result.cases[0].agent_version == "" + + +@pytest.mark.asyncio +async def test_whole_trace_uses_the_last_llm_span_that_has_a_conversation() -> None: + reader: Final = FakeReader( + { + ("t1", "early"): detail("early", "first"), + ("t1", "late"): detail("late", "second"), + ("t1", "tool"): detail("tool", "not an llm"), + ("t1", "text"): detail("text", "raw", input_ui=UIText(kind="text", text="raw")), + }, + traces={ + "t1": trace( + span("early", "llm", 1), span("late", "llm", 5), span("tool", "tool", 9), span("text", "llm", 7) + ) + }, + ) + result: Final = await build(reader, TraceSource(trace_id="t1")) + assert tuple(c.source.span_id for c in result.cases) == ("late",) + assert result.cases[0].messages[-1].content == "second" + + +@pytest.mark.asyncio +async def test_finding_yields_one_case_per_distinct_evidence_span_located_by_its_execution() -> None: + first: Final = execution_id("traces", "alpha", "t1", "ref1") + second: Final = execution_id("traces", "alpha", "t2") + reader: Final = FakeReader( + {("t1", "s1"): detail("s1", "one"), ("t2", "s2"): detail("s2", "two")}, + findings=( + stored_finding("f1", Evidence(execution_id=first, span_id="s1", quote="a")), + stored_finding( + "f2", + Evidence(execution_id=first, span_id="s1", quote="b"), + Evidence(execution_id=second, span_id="s2", quote="c"), + ), + ), + ) + result: Final = await build(reader, FindingSource(lens_id="lens", finding_ids=("f1", "f2"))) + + assert tuple(c.source for c in result.cases) == ( + CaseSource(trace_id="t1", trace_ref="ref1", span_id="s1", finding_id="f1", lens_id="lens"), + CaseSource(trace_id="t2", trace_ref="", span_id="s2", finding_id="f2", lens_id="lens"), + ) + assert result.skipped == () + + +@pytest.mark.asyncio +async def test_text_jsonl_lines_become_cases_and_plain_text_becomes_one_user_message() -> None: + lines: Final = "\n".join( + ( + json.dumps({"messages": [{"role": "user", "content": "hi"}], "reply": "hello", "expected": "greet"}), + "", + json.dumps({"messages": [{"role": "user", "content": "bye"}]}), + ) + ) + jsonl: Final = await build(FakeReader({}), TextSource(text=lines)) + plain: Final = await build(FakeReader({}), TextSource(text="Cancel my order\nplease")) + + assert tuple((c.messages[0].content, c.reply, c.expected) for c in jsonl.cases) == ( + ("hi", "hello", "greet"), + ("bye", "", ""), + ) + assert plain.cases[0].messages == (DatasetMessage(role="user", content="Cancel my order\nplease"),) + + +@pytest.mark.asyncio +async def test_building_again_from_the_same_sources_adds_nothing() -> None: + reader: Final = FakeReader({("t1", "s1"): detail("s1", "one"), ("t1", "s2"): detail("s2", "two")}) + request: Final = BuildRequest( + sources=(TraceSource(trace_id="t1", span_id="s1"), TraceSource(trace_id="t1", span_id="s2")) + ) + first: Final = await build_cases(request, reader, ()) + second: Final = await build_cases(request, reader, first.cases) + + assert len(first.cases) == 2 + assert second.cases == () + assert tuple(s.reason for s in second.skipped) == ("duplicate", "duplicate") + + +@pytest.mark.asyncio +async def test_each_skip_reason_is_reported_against_its_source() -> None: + huge: Final = "x" * (LENS_DATASET_MAX_CASE_CHARS + 1) + reader: Final = FakeReader( + {("t1", "s1"): detail("s1", "same"), ("t1", "big"): detail("big", huge), ("t1", "s3"): detail("s3", "new")} + ) + full: Final = tuple( + DatasetCase(id=str(i), messages=(DatasetMessage(role="user", content=str(i)),), source=CaseSource()) + for i in range(LENS_DATASET_MAX_CASES - 1) + ) + result: Final = await build_cases( + BuildRequest( + sources=( + TraceSource(trace_id="t1", span_id="missing"), + TraceSource(trace_id="t1", span_id="big"), + TraceSource(trace_id="t1", span_id="s1"), + TraceSource(trace_id="t1", span_id="s1"), + TraceSource(trace_id="t1", span_id="s3"), + ) + ), + reader, + full, + ) + + assert tuple(c.source.span_id for c in result.cases) == ("s1",) + assert tuple((s.source.span_id, s.reason) for s in result.skipped) == ( + ("missing", "no_content"), + ("big", "too_large"), + ("s1", "over_limit"), + ("s3", "over_limit"), + ) + + +class CountingReader(FakeReader): + def __init__(self, spans: Mapping[tuple[str, str], SpanDetail]) -> None: + super().__init__(spans) + self.reads: list[str] = [] # mutable-ok: records which spans the build actually fetched + + async def span(self, trace_id: str, span_id: str, trace_ref: str) -> SpanDetail | None: + self.reads.append(span_id) + return await super().span(trace_id, span_id, trace_ref) + + +@pytest.mark.asyncio +async def test_sources_past_the_case_limit_are_not_read() -> None: + reader: Final = CountingReader({("t1", "s1"): detail("s1", "a"), ("t1", "s2"): detail("s2", "b")}) + full: Final = tuple( + DatasetCase(id=str(i), messages=(DatasetMessage(role="user", content=str(i)),), source=CaseSource()) + for i in range(LENS_DATASET_MAX_CASES - 1) + ) + result: Final = await build_cases( + BuildRequest(sources=(TraceSource(trace_id="t1", span_id="s1"), TraceSource(trace_id="t1", span_id="s2"))), + reader, + full, + ) + + assert reader.reads == ["s1"] + assert result.skipped == (SkippedCase(source=CaseSource(trace_id="t1", span_id="s2"), reason="over_limit"),) + + +@pytest.mark.asyncio +async def test_a_malformed_jsonl_line_is_skipped_without_losing_the_valid_lines() -> None: + good: Final = json.dumps({"messages": [{"role": "user", "content": "refund?"}], "reply": "No"}) + result: Final = await build(FakeReader({}), TextSource(text=f'{good}\n{{not json\n{{"reply": "no messages"}}')) + + assert tuple(c.reply for c in result.cases) == ("No",) + assert tuple(s.reason for s in result.skipped) == ("invalid", "invalid") + + +@pytest.mark.asyncio +async def test_an_exported_case_rebuilds_with_its_source_and_agent_version() -> None: + original: Final = DatasetCase( + id="", + messages=(DatasetMessage(role="user", content="refund?"),), + reply="No", + expected="Decline politely", + source=CaseSource(trace_id="t1", span_id="s1"), + agent_version="v7", + ) + result: Final = await build(FakeReader({}), TextSource(text=export_jsonl((original,)))) + + rebuilt: Final = result.cases[0] + assert (rebuilt.source, rebuilt.agent_version, rebuilt.expected) == (original.source, "v7", "Decline politely") + + +@pytest.mark.parametrize( + "case", + ( + DatasetCase( + id="", + messages=(DatasetMessage(role="user", content="q"),), + expected="x" * LENS_DATASET_MAX_CASE_CHARS, + source=CaseSource(), + ), + DatasetCase( + id="", + messages=(DatasetMessage(role="user", content="q", name="x" * LENS_DATASET_MAX_CASE_CHARS),), + source=CaseSource(), + ), + ), +) +def test_expected_and_message_names_count_toward_the_case_size_limit(case: DatasetCase) -> None: + assert revision_problem((case,)) is not None + + +def raw_span(span_id: str, input_ui: UIContent, output_ui: UIContent, raw_input: str, raw_output: str) -> SpanDetail: + return SpanDetail( + span_id=span_id, input_ui=input_ui, output_ui=output_ui, input=raw_input, output=raw_output, attributes={} + ) + + +@pytest.mark.asyncio +async def test_spans_without_chat_messages_fall_back_to_their_text_or_raw_input_and_output() -> None: + fields: Final = UIFields(kind="fields", fields=(UIField(key="q", value="v"),)) + reader: Final = FakeReader( + { + ("t1", "text"): raw_span( + "text", UIText(kind="text", text="shown"), UIText(kind="text", text="answer"), "asked", "raw answer" + ), + ("t1", "fields"): raw_span("fields", fields, fields, '{"q":"v"}', '{"ok":true}'), + } + ) + result: Final = await build( + reader, TraceSource(trace_id="t1", span_id="text"), TraceSource(trace_id="t1", span_id="fields") + ) + + assert tuple((c.messages, c.reply) for c in result.cases) == ( + ((DatasetMessage(role="user", content="asked"),), "answer"), + ((DatasetMessage(role="user", content='{"q":"v"}'),), '{"ok":true}'), + ) + + +@pytest.mark.asyncio +async def test_a_span_with_blank_input_and_output_is_skipped_as_no_content() -> None: + blank: Final = UIText(kind="text", text=" ") + reader: Final = FakeReader({("t1", "s1"): raw_span("s1", blank, blank, " ", "")}) + result: Final = await build(reader, TraceSource(trace_id="t1", span_id="s1")) + + assert result.cases == () + assert result.skipped == (SkippedCase(source=CaseSource(trace_id="t1", span_id="s1"), reason="no_content"),) + + +@pytest.mark.asyncio +async def test_a_whole_trace_without_a_conversation_or_that_is_missing_is_skipped_as_no_content() -> None: + reader: Final = FakeReader( + {("t1", "text"): detail("text", "raw", input_ui=UIText(kind="text", text="raw"))}, + traces={"t1": trace(span("text", "llm", 1), span("gone", "llm", 2))}, + ) + result: Final = await build(reader, TraceSource(trace_id="t1"), TraceSource(trace_id="missing")) + + assert result.cases == () + assert result.skipped == ( + SkippedCase(source=CaseSource(trace_id="t1"), reason="no_content"), + SkippedCase(source=CaseSource(trace_id="missing"), reason="no_content"), + ) + + +@pytest.mark.asyncio +async def test_finding_evidence_with_an_undecodable_execution_or_a_missing_span_is_skipped() -> None: + located: Final = execution_id("traces", "alpha", "t1", "ref1") + reader: Final = FakeReader( + {}, + findings=( + stored_finding( + "f1", + Evidence(execution_id="not-an-execution", span_id="s1", quote="a"), + Evidence(execution_id=located, span_id="gone", quote="b"), + ), + ), + ) + result: Final = await build(reader, FindingSource(lens_id="lens", finding_ids=("f1",))) + + assert result.cases == () + assert result.skipped == ( + SkippedCase(source=CaseSource(span_id="s1", finding_id="f1", lens_id="lens"), reason="no_content"), + SkippedCase( + source=CaseSource(trace_id="t1", trace_ref="ref1", span_id="gone", finding_id="f1", lens_id="lens"), + reason="no_content", + ), + ) + + +def test_export_writes_only_included_cases_one_per_line() -> None: + kept: Final = DatasetCase(id="a", messages=(DatasetMessage(role="user", content="keep"),), source=CaseSource()) + dropped: Final = kept.model_copy(update={"id": "b", "included": False}) + lines: Final = export_jsonl((kept, dropped)).splitlines() + assert tuple(DatasetCase.model_validate_json(line) for line in lines) == (kept,) diff --git a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx index 2c796ebd6fe..1a483f55e90 100644 --- a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx @@ -11,6 +11,7 @@ import { Tabs, TabsContent } from "@/components/ui/tabs"; import { LensServicesProvider, useLensAccessToken, useLensApi, useLiveLensServices } from "./data/LensServices"; import { isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles"; import { InvestigationsView } from "./investigations/InvestigationsView"; +import { DatasetsView } from "./datasets/DatasetsView"; import { LensSettings } from "./settings/LensSettings"; import { createLensDemo } from "./data/demo/createLensDemo"; import { lensQueries } from "./data/queries"; @@ -213,6 +214,9 @@ function LensContent({ userRole, readOnly }: Omit

)} + + + )} {workers && list && ( @@ -237,6 +241,12 @@ function LensContent({ userRole, readOnly }: Omit ); } +function DatasetsPanel({ canView, isAdmin, readOnly }: { canView: boolean; isAdmin: boolean; readOnly: boolean }) { + if (!canView) + return

Datasets require proxy administrator access.

; + return ; +} + function needsSetup( state: LensReadiness, location: { @@ -249,7 +259,7 @@ function needsSetup( issueKey: string | null; }, ) { - if (location.tab === "settings") return false; + if (location.tab === "settings" || location.tab === "datasets") return false; if (location.requested) return true; const selected = location.tab === "traces" ? location.trace : location.lensId || location.dialog || location.issueKey; if (!state.missingTraces || selected) return false; diff --git a/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx b/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx index a4804212012..2895976de01 100644 --- a/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx +++ b/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx @@ -26,6 +26,8 @@ function useLensServices(): LensServices { export const useLensApi = (): LensApi => useLensServices().lens; +export const useOptionalLensApi = (): LensApi | null => useContext(LensServicesContext)?.lens ?? null; + export const useLensAccessToken = (): string => useLensServices().accessToken; export function useLiveLensServices(accessToken: string): LensServices { 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 36c1d9238b2..422601c44c2 100644 --- a/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts +++ b/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts @@ -2,6 +2,7 @@ import { ApiError } from "@/lib/http/client"; import type { TracesApi } from "@/components/lens/traces/api"; import type { LensServices } from "../LensServices"; import type { LensApi } from "../service"; +import { demoDatasetsApi } from "./demoDatasets"; import { createLensDemoData, type LensDemoData } from "./fixtures"; const notInDemo = (): Promise => @@ -14,6 +15,7 @@ function demoLensApi(data: LensDemoData): LensApi { const jobs = (lensId: string) => data.lenses.find((lens) => lens.id === lensId)?.jobs; return { scope: "demo", + datasets: demoDatasetsApi(data), lenses: async () => ({ lenses: data.lenses, workers: [], tracing_enabled: true }), activity: async () => ({ traces: true, requests: false }), runs: (lensId, offset) => found(jobs(lensId)?.slice(offset)), diff --git a/ui/litellm-dashboard/src/components/lens/data/demo/demoDatasets.ts b/ui/litellm-dashboard/src/components/lens/data/demo/demoDatasets.ts new file mode 100644 index 00000000000..d3bd54e3d38 --- /dev/null +++ b/ui/litellm-dashboard/src/components/lens/data/demo/demoDatasets.ts @@ -0,0 +1,206 @@ +import { ApiError } from "@/lib/http/client"; +import type { SpanDetail } from "@/components/lens/traces/types"; + +import { evidenceTarget } from "../../model/findings"; +import type { DatasetsApi } from "../../datasets/client"; +import type { BuildSource, CaseSource, Dataset, DatasetCase, DatasetMessage, SkippedCase } from "../../datasets/types"; +import type { LensDemoData } from "./fixtures"; + +const MAX_CASES = 200; +const NO_SOURCE: CaseSource = { trace_id: "", trace_ref: "", span_id: "", finding_id: "", lens_id: "" }; +const ROLES: ReadonlySet = new Set(["system", "user", "assistant", "tool"]); + +type Candidate = DatasetCase | SkippedCase; + +const isRole = (role: unknown): role is DatasetMessage["role"] => typeof role === "string" && ROLES.has(role); + +function contentHash(text: string): string { + const hash = [...text].reduce((acc, char) => (Math.imul(acc, 31) + char.charCodeAt(0)) | 0, 7); + return (hash >>> 0).toString(16).padStart(8, "0"); +} + +function parsedMessages(raw: string): DatasetMessage[] { + try { + const parsed: unknown = JSON.parse(raw); + if (!Array.isArray(parsed)) return []; + return parsed.flatMap((item: { role?: unknown; content?: unknown }) => + isRole(item?.role) && typeof item.content === "string" + ? [{ role: item.role, content: item.content, name: "", tool_calls: [] }] + : [], + ); + } catch { + return []; + } +} + +function makeCase(messages: DatasetMessage[], reply: string, source: CaseSource): Candidate { + if (messages.length === 0 && !reply) return { source, reason: "no_content" }; + const id = contentHash(JSON.stringify({ messages, reply })); + return { id, messages, reply, tool_calls: [], expected: "", included: true, source, agent_version: "" }; +} + +function caseFromSpan(detail: SpanDetail, source: CaseSource): Candidate { + const messages = parsedMessages(detail.input); + const replies = parsedMessages(detail.output).filter((message) => message.role === "assistant"); + const reply = replies.at(-1)?.content ?? detail.output; + return makeCase( + messages.length ? messages : [{ role: "user", content: detail.input, name: "", tool_calls: [] }], + reply, + source, + ); +} + +function textCases(text: string): Candidate[] { + const lines = text + .split("\n") + .map((line) => line.trim()) + .filter(Boolean); + const parsed = lines.map((line) => { + try { + const value: { messages?: unknown; reply?: unknown } = JSON.parse(line); + return Array.isArray(value.messages) ? value : null; + } catch { + return null; + } + }); + if (parsed.length === 0 || parsed.some((value) => value === null)) + return [makeCase([{ role: "user", content: text, name: "", tool_calls: [] }], "", NO_SOURCE)]; + return parsed.map((value) => + makeCase( + parsedMessages(JSON.stringify(value?.messages)), + typeof value?.reply === "string" ? value.reply : "", + NO_SOURCE, + ), + ); +} + +const isCase = (candidate: Candidate): candidate is DatasetCase => "id" in candidate; + +export function demoDatasetsApi(data: LensDemoData, now: () => Date = () => new Date()): DatasetsApi { + const revisions = new Map(); + const missing = () => Promise.reject(new ApiError("Dataset not found", 404, { detail: "Dataset not found" })); + const latest = (id: string) => revisions.get(id)?.at(-1); + const run = (traceId: string) => data.runs.find(({ trace }) => trace.summary.trace_id === traceId); + const spanCase = (traceId: string, spanId: string, source: CaseSource): Candidate => { + const found = run(traceId); + const detail = spanId + ? found?.details.find((span) => span.span_id === spanId) + : found?.details.findLast((span) => found.trace.spans.find((s) => s.span_id === span.span_id)?.type === "llm"); + return detail ? caseFromSpan(detail, source) : { source, reason: "no_content" }; + }; + const sourceCases = (source: BuildSource): Candidate[] => { + if (source.kind === "text") return textCases(source.text); + if (source.kind === "finding") { + const findings = data.lenses + .flatMap((lens) => lens.jobs.flatMap((job) => job.findings ?? [])) + .filter((finding) => source.finding_ids.includes(finding.id)); + const spans = new Map( + findings.flatMap((finding) => + finding.evidence.map((evidence) => { + const traceId = evidenceTarget(evidence.execution_id)?.id ?? ""; + const origin = { ...NO_SOURCE, trace_id: traceId, span_id: evidence.span_id, finding_id: finding.id }; + return [`${traceId}/${evidence.span_id}`, { ...origin, lens_id: source.lens_id }] as const; + }), + ), + ); + return [...spans.values()].map((origin) => spanCase(origin.trace_id, origin.span_id, origin)); + } + const origin = { + ...NO_SOURCE, + trace_id: source.trace_id, + trace_ref: source.trace_ref ?? "", + span_id: source.span_id ?? "", + }; + return [spanCase(source.trace_id, source.span_id ?? "", origin)]; + }; + return { + list: async () => + [...revisions.values()].map((history) => { + const dataset = history.at(-1)!; + return { + id: dataset.id, + name: dataset.name, + agent_name: dataset.agent_name, + revision: dataset.revision, + case_count: dataset.cases.length, + updated_at: dataset.created_at, + }; + }), + get: (id, revision) => { + const history = revisions.get(id); + const found = revision === undefined ? history?.at(-1) : history?.find((item) => item.revision === revision); + return found ? Promise.resolve(found) : missing(); + }, + create: async ({ name, agent_name }) => { + const dataset: Dataset = { + id: `demo-dataset-${revisions.size + 1}`, + name, + agent_name: agent_name ?? "", + team_id: "", + created_at: now().toISOString(), + revision: 0, + created_by: "demo", + cases: [], + }; + revisions.set(dataset.id, [dataset]); + return dataset; + }, + build: async ({ sources, dataset_id }) => { + const existing = (dataset_id && latest(dataset_id)?.cases) || []; + const seen = new Set(existing.map((item) => item.id)); + const candidates = sources.flatMap(sourceCases); + const cases: DatasetCase[] = []; + const skipped: SkippedCase[] = []; + for (const candidate of candidates) { + if (!isCase(candidate)) skipped.push(candidate); + else if (seen.has(candidate.id)) skipped.push({ source: candidate.source, reason: "duplicate" }); + else if (existing.length + cases.length >= MAX_CASES) + skipped.push({ source: candidate.source, reason: "over_limit" }); + else { + seen.add(candidate.id); + cases.push(candidate); + } + } + return { cases, skipped }; + }, + saveRevision: async (id, { base_revision, cases }) => { + const current = latest(id); + if (!current) return missing(); + if (current.revision !== base_revision) + throw new ApiError("Dataset changed, reload", 409, { detail: "Dataset changed, reload" }); + const saved: Dataset = { + ...current, + revision: current.revision + 1, + created_at: now().toISOString(), + cases: cases.map((item) => ({ + ...item, + reply: item.reply ?? "", + tool_calls: item.tool_calls ?? [], + expected: item.expected ?? "", + included: item.included ?? true, + agent_version: item.agent_version ?? "", + messages: item.messages.map((message) => ({ + ...message, + name: message.name ?? "", + tool_calls: message.tool_calls ?? [], + })), + source: { ...NO_SOURCE, ...item.source }, + })), + }; + revisions.get(id)!.push(saved); + return saved; + }, + exportJsonl: async (id, revision) => { + const history = revisions.get(id); + const found = revision === undefined ? history?.at(-1) : history?.find((item) => item.revision === revision); + if (!found) return missing(); + const lines = found.cases.filter((item) => item.included).map((item) => `${JSON.stringify(item)}\n`); + return new Blob(lines, { type: "application/x-ndjson" }); + }, + evalCases: async (id, revision) => { + const found = revisions.get(id)?.find((item) => item.revision === revision); + if (!found) return missing(); + return { dataset_id: id, revision, cases: found.cases.filter((item) => item.included) }; + }, + }; +} diff --git a/ui/litellm-dashboard/src/components/lens/data/service.ts b/ui/litellm-dashboard/src/components/lens/data/service.ts index f035d882259..9bd7919308b 100644 --- a/ui/litellm-dashboard/src/components/lens/data/service.ts +++ b/ui/litellm-dashboard/src/components/lens/data/service.ts @@ -3,6 +3,7 @@ import type { ApiClient } from "@/lib/http/client"; import { getAuthHeaderName } from "@/lib/http/runtime"; import type { Client } from "openapi-fetch"; import type { components, paths } from "@/lib/http/schema"; +import { liveDatasetsApi, type DatasetsApi } from "../datasets/client"; import type { ActivitySelection, AnalysisModelInfo, @@ -45,6 +46,7 @@ export interface AnalysisKeyRequest { export interface LensApi { /** Partitions query caches between backends (one token, or the demo). */ readonly scope: string; + readonly datasets: DatasetsApi; lenses(): Promise; activity(): Promise; runs(lensId: string, offset: number): Promise; @@ -87,6 +89,7 @@ export function liveLensApi(client: LensClient, apiClient: ApiClient, accessToke const worker = (worker_id: string) => ({ headers, params: { path: { worker_id } } }); return { scope: accessToken, + datasets: liveDatasetsApi(client, apiClient, accessToken), lenses: () => required(client.GET("/lens", { headers })), activity: () => required(client.GET("/lens/activity/available", { headers })), runs: (lensId, offset) => diff --git a/ui/litellm-dashboard/src/components/lens/datasets/AddToDatasetDialog.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/datasets/AddToDatasetDialog.integration.test.tsx new file mode 100644 index 00000000000..c1a06150988 --- /dev/null +++ b/ui/litellm-dashboard/src/components/lens/datasets/AddToDatasetDialog.integration.test.tsx @@ -0,0 +1,138 @@ +import { fireEvent, screen, waitFor, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, expect, it, vi } from "vitest"; + +import { renderWithLens, stubGateway } from "@/../tests/lens-test-utils"; +import { testQueryClient } from "@/../tests/test-utils"; + +import { AddToDatasetDialog } from "./AddToDatasetDialog"; +import type { Dataset, DatasetCase } from "./types"; + +const source = { trace_id: "trace-1", trace_ref: "", span_id: "", finding_id: "", lens_id: "" }; +const makeCase = (id: string, question: string, reply: string): DatasetCase => ({ + id, + messages: [{ role: "user", content: question, name: "", tool_calls: [] }], + reply, + tool_calls: [], + expected: "", + included: true, + source, + agent_version: "", +}); + +const kept = makeCase("kept", "Where is order 1042?", "It ships tomorrow."); +const junk = makeCase("junk", "asdf", "I did not understand."); +const existing = makeCase("existing", "Can I return headphones?", "Yes, within 30 days."); +const dataset: Dataset = { + id: "ds-1", + name: "support cases", + agent_name: "support_agent", + team_id: "", + created_at: "2026-10-01T00:00:00Z", + revision: 3, + created_by: "admin", + cases: [existing], +}; +const sources = [{ kind: "trace" as const, trace_id: "trace-1", trace_ref: "", span_id: "" }]; + +let proxy = stubGateway(); + +beforeEach(() => { + testQueryClient.clear(); + proxy = stubGateway(); + proxy.get.mockImplementation(async (path) => { + if (path === "/lens/datasets") + return [ + { id: "ds-1", name: "support cases", agent_name: "support_agent", revision: 3, case_count: 1, updated_at: "" }, + ]; + if (path === "/lens/datasets/ds-1") return dataset; + throw new Error(`unexpected GET ${path}`); + }); + proxy.post.mockImplementation(async (path) => { + if (path === "/lens/datasets/build") + return { cases: [kept, junk], skipped: [{ source: { ...source, trace_id: "trace-2" }, reason: "duplicate" }] }; + if (path === "/lens/datasets/ds-1/revisions") return { ...dataset, revision: 4 }; + throw new Error(`unexpected POST ${path}`); + }); +}); + +it("saves the dataset's existing cases plus only the ticked new ones, with their expected text", async () => { + const user = userEvent.setup(); + const onClose = vi.fn(); + renderWithLens(); + + const list = await screen.findByRole("list", { name: "Cases to add" }); + expect(within(list).getAllByRole("listitem")[0]).toHaveTextContent("Where is order 1042?It ships tomorrow."); + expect(within(screen.getByRole("list", { name: "Skipped" })).getByText("Already in the dataset")).toBeVisible(); + expect(proxy.post.mock.calls.find(([path]) => path === "/lens/datasets/build")?.[1].body).toEqual({ + sources, + dataset_id: "ds-1", + }); + + await user.click(screen.getByRole("checkbox", { name: "Include Case 2" })); + expect(screen.getByText("1 of 2 selected")).toBeVisible(); + fireEvent.change(screen.getByRole("textbox", { name: "Expected for Case 1" }), { + target: { value: "Gives the ship date" }, + }); + await user.click(screen.getByRole("button", { name: "Save 1 case" })); + + await waitFor(() => expect(onClose).toHaveBeenCalled()); + const saved = proxy.post.mock.calls.find(([path]) => path === "/lens/datasets/ds-1/revisions")?.[1].body; + expect(saved).toEqual({ + base_revision: 3, + cases: [existing, { ...kept, expected: "Gives the ship date" }], + }); +}); + +it("tells the user to reload when someone else saved the dataset first", async () => { + const user = userEvent.setup(); + const onClose = vi.fn(); + vi.stubGlobal( + "fetch", + vi.fn(async (input: RequestInfo | URL, init?: RequestInit) => { + const url = new URL(input instanceof Request ? input.url : String(input), "http://localhost"); + const method = input instanceof Request ? input.method : init?.method ?? "GET"; + if (method === "POST" && url.pathname === "/lens/datasets/ds-1/revisions") + return Response.json({ detail: "Dataset changed, reload" }, { status: 409 }); + if (method === "POST") return Response.json({ cases: [kept], skipped: [] }); + if (url.pathname === "/lens/datasets/ds-1") return Response.json(dataset); + return Response.json([ + { id: "ds-1", name: "support cases", agent_name: "support_agent", revision: 3, case_count: 1, updated_at: "" }, + ]); + }), + ); + renderWithLens(); + + await user.click(await screen.findByRole("button", { name: "Save 1 case" })); + + expect(await screen.findByRole("alert")).toHaveTextContent("Reload to add to the latest version"); + expect(screen.getByRole("button", { name: "Save 1 case" })).toBeDisabled(); + expect(onClose).not.toHaveBeenCalled(); +}); + +it("creates a new dataset by name and saves the cases into its first revision", async () => { + const user = userEvent.setup(); + const created: Dataset = { ...dataset, id: "ds-new", name: "refund cases", revision: 0, cases: [] }; + proxy.get.mockImplementation(async (path) => { + if (path === "/lens/datasets") return []; + throw new Error(`unexpected GET ${path}`); + }); + proxy.post.mockImplementation(async (path) => { + if (path === "/lens/datasets/build") return { cases: [kept], skipped: [] }; + if (path === "/lens/datasets") return created; + if (path === "/lens/datasets/ds-new/revisions") return { ...created, revision: 1 }; + throw new Error(`unexpected POST ${path}`); + }); + renderWithLens(); + + fireEvent.change(await screen.findByRole("textbox", { name: "Name" }), { target: { value: "refund cases" } }); + await user.click(await screen.findByRole("button", { name: "Save 1 case" })); + + await waitFor(() => + expect(proxy.post.mock.calls.map(([path, request]) => [path, request.body])).toEqual([ + ["/lens/datasets/build", { sources, dataset_id: "" }], + ["/lens/datasets", { name: "refund cases", agent_name: "support_agent" }], + ["/lens/datasets/ds-new/revisions", { base_revision: 0, cases: [kept] }], + ]), + ); +}); diff --git a/ui/litellm-dashboard/src/components/lens/datasets/AddToDatasetDialog.tsx b/ui/litellm-dashboard/src/components/lens/datasets/AddToDatasetDialog.tsx new file mode 100644 index 00000000000..e36a8a09e97 --- /dev/null +++ b/ui/litellm-dashboard/src/components/lens/datasets/AddToDatasetDialog.tsx @@ -0,0 +1,424 @@ +"use client"; + +import { useQueryClient } from "@tanstack/react-query"; +import { Check, DatabaseZap, Loader2, RotateCw, TriangleAlert } from "lucide-react"; +import { useId, useState } from "react"; + +import { Button } from "@/components/ui/button"; +import { Checkbox } from "@/components/ui/checkbox"; +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, +} from "@/components/ui/dialog"; +import { Input } from "@/components/ui/input"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Textarea } from "@/components/ui/textarea"; +import { extractProxyErrorMessage } from "@/lib/http/client"; +import { cn } from "@/lib/cva.config"; +import { toast } from "@/lib/toast"; + +import { useTracesLive } from "../traces/api"; +import { useOptionalLensApi } from "../data/LensServices"; +import { useOptionalOnboarding } from "../onboarding/OnboardingContext"; +import { useDatasetRoute, useLensRoute } from "../route"; +import { + datasetKeys, + isRevisionConflict, + useBuildCases, + useCreateDataset, + useDataset, + useDatasets, + useSaveRevision, +} from "./api"; +import { + casePrompt, + caseReplySummary, + EMPTY_DRAFT, + revisionCases, + setExpected, + SKIP_REASON_TEXT, + toggleCase, + type Draft, +} from "./draft"; +import type { BuildSource, Dataset, DatasetCase, DatasetMessage, DatasetSummary, SkippedCase } from "./types"; + +const NEW_DATASET = "__new__"; + +const cases = (count: number): string => `${count} ${count === 1 ? "case" : "cases"}`; + +/** Admins can save cases; the sample session keeps its datasets in memory so anyone can try it. */ +export function useCanAddToDataset(): boolean { + const api = useOptionalLensApi(); + const onboarding = useOptionalOnboarding(); + const live = useTracesLive(); + const canWrite = !!onboarding && onboarding.canInvestigate && !onboarding.readOnly; + return !!api && (!live || canWrite); +} + +export interface AddToDatasetDialogProps { + readonly sources: readonly BuildSource[]; + readonly agentName?: string; + readonly onClose: () => void; +} + +function useOpenSavedDataset() { + const { setTab } = useLensRoute(); + const { openDataset } = useDatasetRoute(); + return (id: string) => { + setTab("datasets"); + openDataset(id); + }; +} + +/** Defaults to the agent's dataset when one exists, otherwise to a new one once the list has loaded. */ +function useDatasetTarget(agentName: string) { + const datasets = useDatasets(); + const [picked, setPicked] = useState(null); + const forAgent = datasets.data?.find((item) => agentName && item.agent_name === agentName)?.id; + const target = picked ?? forAgent ?? (datasets.isPending ? null : NEW_DATASET); + const existingId = target === NEW_DATASET ? null : target; + return { datasets: datasets.data ?? [], target, existingId, setPicked }; +} + +export function AddToDatasetDialog({ sources, agentName = "", onClose }: AddToDatasetDialogProps) { + const { datasets, target, existingId, setPicked } = useDatasetTarget(agentName); + const dataset = useDataset(existingId); + const build = useBuildCases({ sources: [...sources], dataset_id: existingId ?? "" }, target !== null); + const [draft, setDraft] = useState(EMPTY_DRAFT); + const [name, setName] = useState(agentName ? `${agentName} cases` : ""); + const [conflict, setConflict] = useState(false); + const create = useCreateDataset(); + const save = useSaveRevision(); + const queryClient = useQueryClient(); + const openSaved = useOpenSavedDataset(); + + const built = build.data?.cases ?? []; + const chosen = built.filter((item) => !draft.excluded.has(item.id)).length; + const busy = create.isPending || save.isPending; + const targetReady = existingId ? dataset.isSuccess : name.trim().length > 0; + const hasCases = build.isSuccess && chosen > 0; + const idle = !busy && !conflict; + const canSave = hasCases && targetReady && idle; + + const pickTarget = (next: string) => { + setPicked(next); + setConflict(false); + }; + + const reload = async () => { + setConflict(false); + await queryClient.invalidateQueries({ queryKey: datasetKeys.all() }); + }; + + const baseDataset = async (): Promise => { + if (existingId) return dataset.data; + return create.mutateAsync({ name: name.trim(), agent_name: agentName }).catch((error: unknown) => { + toast.fromError(error); + return undefined; + }); + }; + + const submit = async () => { + const base = await baseDataset(); + if (!base) return; + try { + const saved = await save.mutateAsync({ + datasetId: base.id, + body: { base_revision: base.revision, cases: revisionCases(base.cases, built, draft) }, + }); + toast.success(`Added ${cases(chosen)} to ${saved.name}`, { + description: ( + + ), + }); + onClose(); + } catch (error) { + if (isRevisionConflict(error)) setConflict(true); + else toast.fromError(error); + } + }; + + return ( + !open && !busy && onClose()}> + + + Add to dataset + Saves a copy of each conversation, so it stays after the trace expires. + + + setDraft((current) => toggleCase(current, id))} + onExpected={(id, text) => setDraft((current) => setExpected(current, id, text))} + /> + + {conflict ? void reload()} /> : } +
+ + +
+
+
+
+ ); +} + +interface TargetFieldsProps { + readonly datasets: readonly DatasetSummary[]; + readonly target: string | null; + readonly name: string; + readonly onTarget: (target: string) => void; + readonly onName: (name: string) => void; +} + +function TargetFields({ datasets, target, name, onTarget, onName }: TargetFieldsProps) { + const nameId = useId(); + const items = [ + { value: NEW_DATASET, label: "New dataset" }, + ...datasets.map((item) => ({ value: item.id, label: item.name })), + ]; + return ( +
+ + {target === NEW_DATASET && ( + + )} +
+ ); +} + +function ConflictNotice({ onReload }: { onReload: () => void }) { + return ( +

+

+ ); +} + +interface CasePreviewProps { + readonly build: ReturnType; + readonly draft: Draft; + readonly chosen: number; + readonly onToggle: (id: string) => void; + readonly onExpected: (id: string, text: string) => void; +} + +function CasePreview({ build, draft, chosen, onToggle, onExpected }: CasePreviewProps) { + if (build.isPending) + return ( +

+

+ ); + if (build.isError) + return ( +

+

+ ); + const { cases, skipped } = build.data; + return ( +
+
+

Cases

+

+ {chosen} of {cases.length} selected +

+
+ {cases.length === 0 ? ( +

Nothing new to add.

+ ) : ( +
    + {cases.map((item, index) => ( + onToggle(item.id)} + onExpected={(text) => onExpected(item.id, text)} + /> + ))} +
+ )} + {skipped.length > 0 && } +
+ ); +} + +function CaseTile({ + item, + index, + selected, + expected, + onToggle, + onExpected, +}: { + item: DatasetCase; + index: number; + selected: boolean; + expected: string; + onToggle: () => void; + onExpected: (text: string) => void; +}) { + const label = `Case ${index + 1}`; + return ( +
  • + {selected && ( +