From 231a46e40b6d0be83e8287c4917ff3b22f243f05 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 10:04:59 -0700 Subject: [PATCH 01/82] feat(azure): add azure_ai/kimi-k2-thinking from Azure Kimi pricing page (#44382) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 22 +++++++++++++++++++ model_prices_and_context_window.json | 22 +++++++++++++++++++ 2 files changed, 44 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 26977e64e34..ea383ef4c11 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -79845,5 +79845,27 @@ "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false + }, + "azure_ai/kimi-k2-thinking": { + "input_cost_per_token": 6e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 2.5e-06, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/kimi/", + "supported_modalities": [ + "text" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_prompt_caching": false, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_video_input": false, + "supports_vision": false } } diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 26977e64e34..ea383ef4c11 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -79845,5 +79845,27 @@ "supports_tool_choice": true, "supports_vision": false, "supports_web_search": false + }, + "azure_ai/kimi-k2-thinking": { + "input_cost_per_token": 6e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 2.5e-06, + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/kimi/", + "supported_modalities": [ + "text" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_prompt_caching": false, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_video_input": false, + "supports_vision": false } } From b024950353efd55e33082f7f5b267f662478e5bb Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 17:33:01 +0000 Subject: [PATCH 02/82] feat(harness): add Harness.TOOL_LOOP, a minimal in-process tool-calling loop (#44391) * feat(harness): add Harness.TOOL_LOOP, a minimal in-process tool-calling loop Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * refactor(harness): name per-tool spec FunctionTool Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm/__init__.py | 1 + litellm/harness/__init__.py | 4 +- litellm/harness/handlers/__init__.py | 6 +- litellm/harness/handlers/tool_loop_handler.py | 345 +++++++++++++ litellm/harness/options.py | 7 +- litellm/harness/sync.py | 2 +- litellm/harness/types.py | 1 + litellm/llms/base_llm/harness/utils.py | 30 +- .../llms/deepagents/harness/transformation.py | 13 +- litellm/llms/tool_loop/__init__.py | 1 + litellm/llms/tool_loop/harness/__init__.py | 1 + .../llms/tool_loop/harness/transformation.py | 119 +++++ litellm/utils.py | 6 +- .../check_provider_folders_documented.py | 1 + .../handlers/test_tool_loop_handler.py | 476 ++++++++++++++++++ tests/unit/harness/test_init.py | 5 +- tests/unit/harness/test_types.py | 1 + tests/unit/llms/tool_loop/__init__.py | 0 tests/unit/llms/tool_loop/harness/__init__.py | 0 .../tool_loop/harness/test_transformation.py | 154 ++++++ 20 files changed, 1148 insertions(+), 25 deletions(-) create mode 100644 litellm/harness/handlers/tool_loop_handler.py create mode 100644 litellm/llms/tool_loop/__init__.py create mode 100644 litellm/llms/tool_loop/harness/__init__.py create mode 100644 litellm/llms/tool_loop/harness/transformation.py create mode 100644 tests/unit/harness/handlers/test_tool_loop_handler.py create mode 100644 tests/unit/llms/tool_loop/__init__.py create mode 100644 tests/unit/llms/tool_loop/harness/__init__.py create mode 100644 tests/unit/llms/tool_loop/harness/test_transformation.py diff --git a/litellm/__init__.py b/litellm/__init__.py index b0761da7f7c..30a9a80e9b2 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -2302,6 +2302,7 @@ _AGENT_EXPORTS: Final = frozenset( "CodexOptions", "OpenCodeOptions", "DeepAgentsOptions", + "ToolLoopOptions", } ) diff --git a/litellm/harness/__init__.py b/litellm/harness/__init__.py index 79322cdc3e5..213591c5949 100644 --- a/litellm/harness/__init__.py +++ b/litellm/harness/__init__.py @@ -1,4 +1,4 @@ -"""Agent harnesses: run Claude Code, Codex, OpenCode or Deep Agents on any LiteLLM model. +"""Agent harnesses: run Claude Code, Codex, OpenCode, Deep Agents or Tool Loop on any LiteLLM model. The entrypoints live on the top-level package: @@ -30,6 +30,7 @@ from litellm.harness.options import ( CodexOptions, DeepAgentsOptions, OpenCodeOptions, + ToolLoopOptions, ) from litellm.harness.runtime import ( AsyncEventStream, @@ -86,6 +87,7 @@ __all__ = ( "StateIncompatible", "Text", "ToolCall", + "ToolLoopOptions", "ToolResult", "Usage", "aagent", diff --git a/litellm/harness/handlers/__init__.py b/litellm/harness/handlers/__init__.py index 4b32d7ca649..5b374a942fa 100644 --- a/litellm/harness/handlers/__init__.py +++ b/litellm/harness/handlers/__init__.py @@ -1,4 +1,4 @@ -"""Handlers run a harness config: CLI runtimes as subprocesses, Deep Agents in-process.""" +"""Handlers run a harness config: CLI runtimes and in-process harnesses.""" from __future__ import annotations @@ -29,6 +29,10 @@ def get_harness_handler(config: BaseHarnessConfig) -> BaseHarnessHandler: from litellm.harness.handlers.deepagents_handler import DeepAgentsHandler return DeepAgentsHandler(config) + if config.harness is Harness.TOOL_LOOP: + from litellm.harness.handlers.tool_loop_handler import ToolLoopHandler + + return ToolLoopHandler(config) raise HarnessError(f"No handler for Harness.{config.harness.name}") diff --git a/litellm/harness/handlers/tool_loop_handler.py b/litellm/harness/handlers/tool_loop_handler.py new file mode 100644 index 00000000000..33efcba4ff7 --- /dev/null +++ b/litellm/harness/handlers/tool_loop_handler.py @@ -0,0 +1,345 @@ +"""In-process tool-calling loop handler.""" + +from __future__ import annotations + +import asyncio +import copy +import inspect +import json +import math +from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Mapping +from dataclasses import dataclass +from types import MappingProxyType +from typing import Final, Protocol, TypeAlias + +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError + +import litellm +from litellm.harness.context import SessionContext +from litellm.harness.errors import CapabilityUnsupported +from litellm.harness.handlers.base import BaseHarnessHandler +from litellm.harness.types import Approval, Event, Reasoning, Text, ToolCall, ToolResult +from litellm.llms.base_llm.harness.transformation import HarnessTurnError +from litellm.llms.tool_loop.harness.transformation import ( + TOOL_LOOP_MAX_MODEL_CALLS, + FunctionTool, + ToolLoopHarnessConfig, + completion_kwargs, + function_tool, +) +from litellm.types.completion import ChatCompletionMessageParam +from litellm.types.utils import ( + ChatCompletionMessageCustomToolCall, + ChatCompletionMessageToolCall, + ChatCompletionToolParam, + ModelResponse, +) + + +@dataclass(frozen=True, slots=True) +class _FunctionToolCall: + id: str + name: str + arguments: str + + +AsyncCompletion: TypeAlias = Callable[..., Awaitable[ModelResponse]] +_ARGUMENTS_ADAPTER: Final = TypeAdapter(dict[str, object]) +_MAPPING_ADAPTER: Final = TypeAdapter(Mapping[str, object]) +_MODEL_RESPONSE_ADAPTER: Final = TypeAdapter(ModelResponse) +_RAW_DECODE_ADAPTER: Final = TypeAdapter(tuple[object, int]) +_JSON_DECODER: Final = json.JSONDecoder() +_HISTORY_ADAPTER: Final = TypeAdapter(list[dict[str, object]]) + + +class _Usage(BaseModel): + model_config = ConfigDict(from_attributes=True) + + prompt_tokens: int | None = None + completion_tokens: int | None = None + + +_USAGE_ADAPTER: Final = TypeAdapter(_Usage) + + +class _AwaitableObject(Protocol): + def __await__(self) -> Generator[object, None, object]: ... + + +async def _await_tool_result(result: _AwaitableObject) -> object: + return await result + + +@dataclass(frozen=True, slots=True) +class _ToolOutcome: + output: str + is_error: bool + + +def _normalize_tool_call( + call: ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall, +) -> _FunctionToolCall: + if isinstance(call, ChatCompletionMessageToolCall): + return _FunctionToolCall( + id=call.id, + name=call.function.name or "", + arguments=call.function.arguments, + ) + return _FunctionToolCall(id=call.id, name=call.custom.name, arguments=call.custom.input) + + +def _parse_arguments(raw: str) -> tuple[dict[str, object], str | None]: # mutable-ok: event input uses a dict + text: Final = raw.strip() + try: + raw_decoded: object = _JSON_DECODER.raw_decode(text) + except json.JSONDecodeError as error: + return {}, f"{type(error).__name__}: {error}" + parsed, end = _RAW_DECODE_ADAPTER.validate_python(raw_decoded) + if text[end:].strip(): + return {}, "JSONDecodeError: Extra data after tool arguments" + try: + arguments: Final = _ARGUMENTS_ADAPTER.validate_python(parsed) + except ValidationError as error: + if not isinstance(parsed, dict): + return {}, "ValueError: tool arguments must be a JSON object" + return {}, f"{type(error).__name__}: {error}" + return arguments, None + + +def _cost_value(value: object) -> float | None: + if isinstance(value, bool): + return None + if not isinstance(value, str | int | float): + return None + try: + cost: Final = float(value) + except (OverflowError, TypeError, ValueError): + return None + return cost if math.isfinite(cost) else None + + +def _as_mapping(value: object) -> Mapping[str, object] | None: + try: + return _MAPPING_ADAPTER.validate_python(value) + except ValidationError: + return None + + +def _response_cost(response: ModelResponse) -> float: + hidden_params_value: Final[object] = getattr(response, "_hidden_params", {}) + hidden_params: Final = _as_mapping(hidden_params_value) + if hidden_params is not None: + additional_headers: Final = _as_mapping(hidden_params.get("additional_headers")) + if additional_headers is not None: + header_cost: Final = _cost_value(additional_headers.get("llm_provider-x-litellm-response-cost")) + if header_cost is not None: + return header_cost + hidden_cost: Final = _cost_value(hidden_params.get("response_cost")) + if hidden_cost is not None: + return hidden_cost + try: + calculated_cost: Final = litellm.completion_cost(completion_response=response) + except Exception: + return 0.0 + return _cost_value(calculated_cost) or 0.0 + + +async def _approval_error(approval: Approval | None) -> str | None: + if approval is None: + return None + allowed, reason = await approval.wait() + return None if allowed else f"denied: {reason}" + + +async def _tool_outcome( + tool: FunctionTool | None, + tool_name: str, + arguments: dict[str, object], + parse_error: str | None, + approval_error: str | None, +) -> _ToolOutcome: + if parse_error is not None: + return _ToolOutcome(output=parse_error, is_error=True) + if approval_error is not None: + return _ToolOutcome(output=approval_error, is_error=True) + if tool is None: + return _ToolOutcome(output=f"ValueError: unknown tool {tool_name!r}", is_error=True) + try: + result: Final = await _execute_tool(tool, arguments) + output: Final = result if isinstance(result, str) else json.dumps(result, default=str) + return _ToolOutcome(output=output, is_error=False) + except Exception as error: + return _ToolOutcome(output=f"{type(error).__name__}: {error}", is_error=True) + + +def _record_usage(ctx: SessionContext, response: ModelResponse) -> None: + try: + usage_value: Final[object] = getattr(response, "usage", None) + usage: Final = _USAGE_ADAPTER.validate_python(usage_value) + input_tokens: Final = usage.prompt_tokens or 0 + output_tokens: Final = usage.completion_tokens or 0 + ctx.calls += 1 # rebind-ok: SessionContext is the runtime's per-session usage sink + ctx.input_tokens += input_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink + ctx.output_tokens += output_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink + ctx.cost += _response_cost(response) # rebind-ok: SessionContext is the runtime's per-session usage sink + except Exception: + return + + +async def _execute_tool(tool: FunctionTool, arguments: Mapping[str, object]) -> object: + validated_model: Final = tool.args_model.model_validate(arguments) + values_object: Final[object] = validated_model.model_dump() + validated: Final = _MAPPING_ADAPTER.validate_python(values_object) + parameters: Final = tuple(inspect.signature(tool.fn).parameters.values()) + positional_args: Final = tuple( + validated[parameter.name] for parameter in parameters if parameter.kind is inspect.Parameter.POSITIONAL_ONLY + ) + keyword_args: Final[dict[str, object]] = { # mutable-ok: tool calls need keyword arguments + parameter.name: validated[parameter.name] + for parameter in parameters + if parameter.kind is not inspect.Parameter.POSITIONAL_ONLY + } + if inspect.iscoroutinefunction(tool.fn): + async_result: Final[object] = tool.fn(*positional_args, **keyword_args) + if inspect.isawaitable(async_result): + return await _await_tool_result(async_result) + return async_result + sync_result: Final[object] = await asyncio.to_thread(tool.fn, *positional_args, **keyword_args) + if inspect.isawaitable(sync_result): + return await _await_tool_result(sync_result) + return sync_result + + +async def _default_acompletion(**kwargs: object) -> ModelResponse: # kwargs-ok: provider-specific completion options + response: Final[object] = await litellm.acompletion(**kwargs) + return _MODEL_RESPONSE_ADAPTER.validate_python(response) + + +class ToolLoopHandler(BaseHarnessHandler): + def __init__( + self, + config: ToolLoopHarnessConfig, + acompletion: AsyncCompletion | None = None, + ) -> None: + super().__init__(config) # pyright: ignore[reportUnknownMemberType] # base handler config is unparameterized + self._config = config + self._acompletion = acompletion if acompletion is not None else _default_acompletion + self._messages: tuple[ChatCompletionMessageParam, ...] = () + self._tools: Mapping[str, FunctionTool] = MappingProxyType({}) + self._tool_specs: tuple[ChatCompletionToolParam, ...] = () + self._completion_kwargs: Mapping[str, object] = MappingProxyType({}) + + async def start(self, ctx: SessionContext) -> None: + self._config.validate_environment(ctx) + tools: Final = tuple(function_tool(fn) for fn in ctx.tools) + if len({tool.name for tool in tools}) != len(tools): + raise ValueError("Harness.TOOL_LOOP tool names must be unique") + self._tools = MappingProxyType({tool.name: tool for tool in tools}) + self._tool_specs = tuple(tool.spec for tool in tools) + self._completion_kwargs = MappingProxyType(completion_kwargs(ctx)) + if not self._messages and ctx.instructions: + self._messages = ({"role": "system", "content": ctx.instructions},) + + def native_session_id(self) -> str | None: + return None + + async def resume(self, ctx: SessionContext, native_session_id: str) -> None: + raise CapabilityUnsupported("Harness.TOOL_LOOP does not support resume") + + async def history( + self, ctx: SessionContext + ) -> list[dict[str, object]]: # mutable-ok: public API returns copied message dictionaries + history_object: Final[object] = copy.deepcopy(list(self._messages)) + return _HISTORY_ADAPTER.validate_python(history_object) # pyright: ignore[reportIncompatibleMethodOverride] # base history uses Any + + async def stop(self, ctx: SessionContext) -> None: + return None + + async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]: + ctx.final_text = "" # rebind-ok: SessionContext is the runtime's per-turn result sink + ctx.output_json = None # rebind-ok: SessionContext is the runtime's per-turn result sink + user_message: Final[ChatCompletionMessageParam] = {"role": "user", "content": prompt} + self._messages = (*self._messages, user_message) + for _ in range(TOOL_LOOP_MAX_MODEL_CALLS): + messages: list[ChatCompletionMessageParam] = copy.deepcopy( # mutable-ok: acompletion takes list messages + list(self._messages) + ) + tool_specs: list[ChatCompletionToolParam] = copy.deepcopy( # mutable-ok: acompletion takes tool list + list(self._tool_specs) + ) + request_kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments + key: value for key, value in self._completion_kwargs.items() if key not in {"messages", "tools"} + } + kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments + **request_kwargs, + "messages": messages, + **({"tools": tool_specs} if tool_specs else {}), + } + response = await self._acompletion(**kwargs) + _record_usage(ctx, response) + message = response.choices[0].message + reasoning_value: object = getattr(message, "reasoning_content", None) + reasoning = reasoning_value if isinstance(reasoning_value, str) else None + content = message.content + if reasoning: + yield Reasoning(reasoning) + if content: + yield Text(content) + tool_calls = message.tool_calls or () + if not tool_calls: + final_text = content or "" + ctx.final_text = final_text # rebind-ok: SessionContext is the runtime's per-turn result sink + ctx.output_json = content if ctx.output is not None else None # rebind-ok: per-turn output sink + final_message: ChatCompletionMessageParam = { + "role": "assistant", + "content": content, + } + self._messages = (*self._messages, final_message) + return + normalized_calls = tuple(_normalize_tool_call(call) for call in tool_calls) + assistant_message: ChatCompletionMessageParam = { + "role": "assistant", + "content": content, + "tool_calls": [ + { + "id": call.id, + "type": "function", + "function": {"name": call.name, "arguments": call.arguments}, + } + for call in normalized_calls + ], + } + self._messages = (*self._messages, assistant_message) + for call in normalized_calls: + arguments, parse_error = _parse_arguments(call.arguments) + yield ToolCall( + id=call.id, + name=call.name, + native_name=call.name, + input=arguments, + builtin=False, + ) + approval = ( + Approval(tool=call.name, input=arguments) + if parse_error is None and ctx.permissions == "ask" + else None + ) + if approval is not None: + yield approval + approval_error = await _approval_error(approval) + tool = self._tools.get(call.name) + outcome = await _tool_outcome( + tool, + call.name, + arguments, + parse_error, + approval_error, + ) + yield ToolResult(id=call.id, output=outcome.output, is_error=outcome.is_error) + tool_message: ChatCompletionMessageParam = { + "role": "tool", + "tool_call_id": call.id, + "content": outcome.output, + } + self._messages = (*self._messages, tool_message) + raise HarnessTurnError(f"Harness.TOOL_LOOP exceeded {TOOL_LOOP_MAX_MODEL_CALLS} model calls in one turn") diff --git a/litellm/harness/options.py b/litellm/harness/options.py index 2e3ec5d0fd5..b0014fea393 100644 --- a/litellm/harness/options.py +++ b/litellm/harness/options.py @@ -34,4 +34,9 @@ class DeepAgentsOptions: recursion_limit: int | None = None -HarnessOptions = ClaudeCodeOptions | CodexOptions | OpenCodeOptions | DeepAgentsOptions +@dataclass(frozen=True) +class ToolLoopOptions: + completion_kwargs: Mapping[str, Any] = field(default_factory=dict) + + +HarnessOptions = ClaudeCodeOptions | CodexOptions | OpenCodeOptions | DeepAgentsOptions | ToolLoopOptions diff --git a/litellm/harness/sync.py b/litellm/harness/sync.py index 1cb8fbcbca2..cdcf251af51 100644 --- a/litellm/harness/sync.py +++ b/litellm/harness/sync.py @@ -411,7 +411,7 @@ def agent( options: HarnessOptions | None = None, install: bool = False, ) -> Result | EventStream: - """Run an agent harness (Claude Code, Codex, OpenCode, Deep Agents) on one prompt. + """Run an agent harness (Claude Code, Codex, OpenCode, Deep Agents, Tool Loop) on one prompt. Returns a Result. With stream=True it returns an iterator of events instead. Prefix the model with `litellm_proxy/` to route every model call through your diff --git a/litellm/harness/types.py b/litellm/harness/types.py index 88ad7f711ea..475568a9979 100644 --- a/litellm/harness/types.py +++ b/litellm/harness/types.py @@ -27,6 +27,7 @@ class Harness(Enum): CODEX = "codex" OPENCODE = "opencode" DEEPAGENTS = "deepagents" + TOOL_LOOP = "tool_loop" def require_harness(harness: object) -> Harness: diff --git a/litellm/llms/base_llm/harness/utils.py b/litellm/llms/base_llm/harness/utils.py index 7c0278d5e36..8dcf784c90f 100644 --- a/litellm/llms/base_llm/harness/utils.py +++ b/litellm/llms/base_llm/harness/utils.py @@ -5,17 +5,33 @@ from __future__ import annotations import itertools import json import os +import typing from collections.abc import Iterator, Mapping, Sequence from types import MappingProxyType -from typing import Any, Final, TypeAlias +from typing import Final, TypeAlias + +from pydantic import TypeAdapter, ValidationError from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH +if typing.TYPE_CHECKING: + from litellm.harness.context import SessionContext + # A decoded JSON document: what json.loads / model_json_schema() produce. JSONValue: TypeAlias = "dict[str, JSONValue] | list[JSONValue] | str | int | float | bool | None" SKILL_MANIFEST: Final = "SKILL.md" _JSON_DECODER: Final = json.JSONDecoder() +_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object]) +_RAW_DECODE_ADAPTER: Final = TypeAdapter(tuple[object, int]) + + +def gateway_headers(ctx: SessionContext) -> dict[str, str]: # mutable-ok: acompletion(extra_headers=) requires dict + metadata_json: Final = json.dumps(dict(ctx.metadata), default=str) if ctx.metadata else None + return { + "x-litellm-tags": f"harness,{ctx.harness.value}", + **({"x-litellm-spend-logs-metadata": metadata_json} if metadata_json is not None else {}), + } def normalize_tool_name(native_name: str, mapping: Mapping[str, str]) -> str: @@ -35,17 +51,18 @@ def last_json_object(text: str) -> str | None: index = text.find("{") while index != -1: try: - obj, end = _JSON_DECODER.raw_decode(text, index) + raw_decoded: object = _JSON_DECODER.raw_decode(text, index) except json.JSONDecodeError: index = text.find("{", index + 1) continue + obj, end = _RAW_DECODE_ADAPTER.validate_python(raw_decoded) if isinstance(obj, dict): last = json.dumps(obj) index = text.find("{", end) return last -def structured_output_instruction(schema: Mapping[str, Any]) -> str: +def structured_output_instruction(schema: Mapping[str, object]) -> str: return ( "When you have finished, your final message must be a single JSON object that " "matches this JSON schema, with no other text before or after it:\n" @@ -82,16 +99,15 @@ def strict_json_schema(schema: JSONValue, depth: int = 0) -> JSONValue: return result -def decode_json_line(line: bytes | str) -> Mapping[str, Any] | None: +def decode_json_line(line: bytes | str) -> Mapping[str, object] | None: """One JSONL line as a dict, or None for blank / non-JSON / non-object lines.""" text = line.strip() if not text: return None try: - obj = json.loads(text) - except json.JSONDecodeError: + return _JSON_OBJECT_ADAPTER.validate_json(text) + except ValidationError: return None - return obj if isinstance(obj, dict) else None def stderr_tail_text(stderr_tail: Sequence[str]) -> str: diff --git a/litellm/llms/deepagents/harness/transformation.py b/litellm/llms/deepagents/harness/transformation.py index 82b2c044eac..14f87f0753b 100644 --- a/litellm/llms/deepagents/harness/transformation.py +++ b/litellm/llms/deepagents/harness/transformation.py @@ -27,6 +27,7 @@ from litellm.harness.types import ( ToolResult, ) from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig +from litellm.llms.base_llm.harness.utils import gateway_headers if TYPE_CHECKING: from litellm.harness.context import SessionContext @@ -70,18 +71,6 @@ APPROVAL_TOOLS: Final = WRITE_TOOLS | EXECUTE_TOOLS _APPROVAL_DECISIONS: Final = ("approve", "reject") -def gateway_headers( - ctx: SessionContext, -) -> dict[str, str]: # mutable-ok: ChatLiteLLM.extra_headers is a pydantic dict field - """Same attribution headers the session endpoint adds for CLI harnesses.""" - metadata = ctx.metadata - metadata_json = json.dumps(dict(metadata), default=str) if metadata else None # mutable-ok: for json.dumps - metadata_header = (("x-litellm-spend-logs-metadata", metadata_json),) if metadata_json is not None else () - return dict( # mutable-ok: ChatLiteLLM.extra_headers is a pydantic dict field - (("x-litellm-tags", f"harness,{ctx.harness.value}"), *metadata_header) - ) - - def chat_model_kwargs( ctx: SessionContext, ) -> dict[str, Any]: # mutable-ok: ChatLiteLLM constructor kwargs, splatted as **kwargs diff --git a/litellm/llms/tool_loop/__init__.py b/litellm/llms/tool_loop/__init__.py new file mode 100644 index 00000000000..f163f698841 --- /dev/null +++ b/litellm/llms/tool_loop/__init__.py @@ -0,0 +1 @@ +"""In-process tool loop harness.""" diff --git a/litellm/llms/tool_loop/harness/__init__.py b/litellm/llms/tool_loop/harness/__init__.py new file mode 100644 index 00000000000..49ca50b61c2 --- /dev/null +++ b/litellm/llms/tool_loop/harness/__init__.py @@ -0,0 +1 @@ +"""Tool Loop harness configuration.""" diff --git a/litellm/llms/tool_loop/harness/transformation.py b/litellm/llms/tool_loop/harness/transformation.py new file mode 100644 index 00000000000..16461916f1a --- /dev/null +++ b/litellm/llms/tool_loop/harness/transformation.py @@ -0,0 +1,119 @@ +"""Configuration and tool-schema helpers for the in-process Tool Loop harness.""" + +from __future__ import annotations + +import inspect +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from types import MappingProxyType +from typing import Final + +from pydantic import BaseModel, ConfigDict, TypeAdapter, create_model + +from litellm.harness.context import SessionContext +from litellm.harness.options import ToolLoopOptions +from litellm.harness.types import Capabilities, Harness +from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig +from litellm.llms.base_llm.harness.utils import gateway_headers +from litellm.types.utils import ChatCompletionToolParam + +TOOL_LOOP_MAX_MODEL_CALLS: Final = 100 +_ANNOTATIONS_ADAPTER: Final = TypeAdapter(Mapping[str, object]) +_OBJECT_ADAPTER: Final = TypeAdapter(object) +_MODEL_FACTORY: Final[Callable[..., type[BaseModel]]] = create_model + + +@dataclass(frozen=True, slots=True) +class FunctionTool: + name: str + fn: Callable[..., object] + args_model: type[BaseModel] + spec: ChatCompletionToolParam + + +def _field_definition( + parameter: inspect.Parameter, + annotations: Mapping[str, object], +) -> tuple[object, object]: + annotation: Final = annotations.get(parameter.name, object) + raw_default: Final[object] = parameter.default # pyright: ignore[reportAny] # inspect exposes defaults as Any + if raw_default is inspect.Parameter.empty: + return annotation, ... + default: Final = _OBJECT_ADAPTER.validate_python(raw_default) + return annotation, default + + +def function_tool(fn: Callable[..., object]) -> FunctionTool: + signature: Final = inspect.signature(fn) + parameters: Final = tuple(signature.parameters.values()) + raw_annotations: Final[object] = inspect.get_annotations(fn, eval_str=True) + annotations: Final = _ANNOTATIONS_ADAPTER.validate_python(raw_annotations) + if any( + parameter.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD) for parameter in parameters + ): + raise ValueError(f"Tool {fn.__name__} cannot use variadic parameters") + + fields: Final = MappingProxyType( + {parameter.name: _field_definition(parameter, annotations) for parameter in parameters} + ) + args_model: Final[type[BaseModel]] = _MODEL_FACTORY( + f"{fn.__name__}_args", + __config__=ConfigDict(extra="forbid"), + **fields, # pyright: ignore[reportCallIssue, reportArgumentType] # Pydantic creates fields dynamically + ) + spec: Final[ChatCompletionToolParam] = { + "type": "function", + "function": { + "name": fn.__name__, + "description": inspect.getdoc(fn) or "", + "parameters": args_model.model_json_schema(), + }, + } + return FunctionTool(name=fn.__name__, fn=fn, args_model=args_model, spec=spec) + + +def _routing_kwargs(ctx: SessionContext) -> Mapping[str, object]: + if not ctx.model: + raise ValueError("Harness.TOOL_LOOP needs model=") + if ctx.gateway is not None: + return { + "model": f"litellm_proxy/{ctx.model}", + "api_base": ctx.gateway.api_base, + "api_key": ctx.gateway.api_key, + "extra_headers": gateway_headers(ctx), + } + return { + "model": ctx.model, + **({"api_key": ctx.api_key} if ctx.api_key is not None else {}), + **({"api_base": ctx.api_base} if ctx.api_base is not None else {}), + } + + +def completion_kwargs(ctx: SessionContext) -> Mapping[str, object]: + options: Final = ToolLoopHarnessConfig().get_options(ctx) + routing: Final = _routing_kwargs(ctx) + kwargs: Final[Mapping[str, object]] = MappingProxyType({**options.completion_kwargs, **routing}) + if ctx.output is None: + return kwargs + return {**kwargs, "response_format": ctx.output} + + +class ToolLoopHarnessConfig(BaseHarnessConfig[ToolLoopOptions]): + harness = Harness.TOOL_LOOP + options_type = ToolLoopOptions + uses_model_endpoint = False + capabilities = Capabilities( + structured_output=True, + tool_approval=True, + tool_filtering=False, + history=True, + custom_tools=True, + skills=False, + resume=False, + permission_modes=frozenset({"ask", "full"}), + ) + + def validate_environment(self, ctx: SessionContext) -> None: + self.get_options(ctx) + if not ctx.model: + raise ValueError("Harness.TOOL_LOOP needs model=") diff --git a/litellm/utils.py b/litellm/utils.py index 9200844a2e3..f80a2851ba6 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9909,7 +9909,7 @@ class ProviderConfigManager: @staticmethod def get_provider_harness_config(harness: Harness) -> BaseHarnessConfig | None: """ - Get the agent-harness configuration (Claude Code, Codex, OpenCode, Deep Agents). + Get the agent-harness configuration (Claude Code, Codex, OpenCode, Deep Agents, Tool Loop). """ from litellm.harness.types import Harness as _Harness @@ -9935,6 +9935,10 @@ class ProviderConfigManager: ) return DeepAgentsHarnessConfig() + if harness == _Harness.TOOL_LOOP: + from litellm.llms.tool_loop.harness.transformation import ToolLoopHarnessConfig + + return ToolLoopHarnessConfig() return None @staticmethod diff --git a/tests/code_coverage_tests/check_provider_folders_documented.py b/tests/code_coverage_tests/check_provider_folders_documented.py index 08fbde3d979..6f58b393cd1 100644 --- a/tests/code_coverage_tests/check_provider_folders_documented.py +++ b/tests/code_coverage_tests/check_provider_folders_documented.py @@ -33,6 +33,7 @@ EXCLUDED_FOLDERS = { "codex", "opencode", "deepagents", + "tool_loop", } diff --git a/tests/unit/harness/handlers/test_tool_loop_handler.py b/tests/unit/harness/handlers/test_tool_loop_handler.py new file mode 100644 index 00000000000..2ca265cda67 --- /dev/null +++ b/tests/unit/harness/handlers/test_tool_loop_handler.py @@ -0,0 +1,476 @@ +"""Tests for the in-process Tool Loop handler.""" + +from __future__ import annotations + +from collections.abc import Callable, Iterator, Mapping +from pathlib import Path +from typing import Final, Literal + +import litellm +import pytest +from pydantic import BaseModel + +from litellm import sandbox +from litellm.harness.context import GatewayTarget, SessionContext +from litellm.harness.handlers.tool_loop_handler import ToolLoopHandler +from litellm.harness.options import ToolLoopOptions +from litellm.harness.types import Approval, Event, Harness, Text, ToolCall, ToolResult +from litellm.llms.base_llm.harness.transformation import HarnessTurnError +from litellm.llms.tool_loop.harness.transformation import ToolLoopHarnessConfig +from litellm.types.utils import ModelResponse + + +class ScriptedCompletion: + def __init__(self, responses: tuple[ModelResponse, ...]) -> None: + self.responses: Iterator[ModelResponse] = iter(responses) + self.calls: list[dict[str, object]] = [] # mutable-ok: captures injected completion requests + + async def __call__(self, **kwargs: object) -> ModelResponse: + self.calls.append(dict(kwargs)) + return next(self.responses) + + +def model_response( + *, + content: str | None = None, + tool_calls: tuple[dict[str, object], ...] = (), + prompt_tokens: int = 0, + completion_tokens: int = 0, + hidden_params: dict[str, object] | None = None, +) -> ModelResponse: + response: Final = ModelResponse( + model="gpt-test", + choices=[ + { + "message": { + "role": "assistant", + "content": content, + "tool_calls": list(tool_calls) if tool_calls else None, + } + } + ], + usage={ + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + "total_tokens": prompt_tokens + completion_tokens, + }, + ) + if hidden_params is not None: + response._hidden_params = hidden_params + return response + + +def function_call(name: str, arguments: str, call_id: str = "call-1") -> dict[str, object]: + return { + "id": call_id, + "type": "function", + "function": {"name": name, "arguments": arguments}, + } + + +def make_context( + tmp_path: Path, + *, + model: str | None = "gpt-4o-mini", + gateway: GatewayTarget | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: tuple[Callable[..., object], ...] = (), + permissions: Literal["ask", "full"] = "full", + output: type[BaseModel] | None = None, + metadata: Mapping[str, object] | None = None, + options: ToolLoopOptions | None = None, +) -> SessionContext: + return SessionContext( + harness=Harness.TOOL_LOOP, + sandbox=sandbox.local(tmp_path), + session_id="tool-loop-session", + model=model, + gateway=gateway, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + permissions=permissions, + output=output, + metadata={} if metadata is None else metadata, + options=options, + ) + + +def make_handler(completion: ScriptedCompletion) -> ToolLoopHandler: + return ToolLoopHandler(ToolLoopHarnessConfig(), acompletion=completion) + + +async def run_turn( + handler: ToolLoopHandler, + ctx: SessionContext, + prompt: str, + allow: bool = True, +) -> tuple[Event, ...]: + events: list[Event] = [] # mutable-ok: gathers this async turn for assertions + async for event in handler.turn(ctx, prompt): + events.append(event) + if isinstance(event, Approval): + event.allow() if allow else event.deny("not approved") + return tuple(events) + + +def add(a: int, b: int) -> int: + return a + b + + +def duplicate_add() -> Callable[..., object]: + def add(a: int, b: int) -> int: + return a + b + + return add + + +async def test_tool_round_trip_appends_assistant_and_tool_messages(tmp_path: Path) -> None: + completion: Final = ScriptedCompletion( + ( + model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)), + model_response(content="The sum is 5"), + ) + ) + ctx: Final = make_context(tmp_path, instructions="Use tools when needed", tools=(add,)) + handler: Final = make_handler(completion) + await handler.start(ctx) + + events: Final = await run_turn(handler, ctx, "Add two and three") + + assert events == ( + ToolCall( + id="call-1", + name="add", + native_name="add", + input={"a": 2, "b": 3}, + builtin=False, + ), + ToolResult(id="call-1", output="5", is_error=False), + Text(delta="The sum is 5"), + ) + assert ctx.final_text == "The sum is 5" + assert completion.calls[1]["messages"] == [ + {"role": "system", "content": "Use tools when needed"}, + {"role": "user", "content": "Add two and three"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "add", "arguments": '{"a": 2, "b": 3}'}, + } + ], + }, + {"role": "tool", "tool_call_id": "call-1", "content": "5"}, + ] + + +async def test_multiple_tool_calls_run_in_order(tmp_path: Path) -> None: + values: list[int] = [] # mutable-ok: records the order of calls from the injected model response + + def record(value: int) -> int: + values.append(value) + return value + + completion: Final = ScriptedCompletion( + ( + model_response( + tool_calls=( + function_call("record", '{"value": 1}', "call-1"), + function_call("record", '{"value": 2}', "call-2"), + ) + ), + model_response(content="Recorded"), + ) + ) + ctx: Final = make_context(tmp_path, tools=(record,)) + handler: Final = make_handler(completion) + await handler.start(ctx) + + events: Final = await run_turn(handler, ctx, "Record both") + + assert values == [1, 2] + assert tuple(event for event in events if isinstance(event, ToolResult)) == ( + ToolResult(id="call-1", output="1", is_error=False), + ToolResult(id="call-2", output="2", is_error=False), + ) + + +async def test_tool_exception_is_returned_to_model_and_loop_continues(tmp_path: Path) -> None: + def fail() -> str: + raise RuntimeError("tool failed") + + completion: Final = ScriptedCompletion( + ( + model_response(tool_calls=(function_call("fail", "{}"),)), + model_response(content="Recovered"), + ) + ) + ctx: Final = make_context(tmp_path, tools=(fail,)) + handler: Final = make_handler(completion) + await handler.start(ctx) + + events: Final = await run_turn(handler, ctx, "Run fail") + + assert ToolResult(id="call-1", output="RuntimeError: tool failed", is_error=True) in events + assert completion.calls[1]["messages"][-1] == { + "role": "tool", + "tool_call_id": "call-1", + "content": "RuntimeError: tool failed", + } + assert ctx.final_text == "Recovered" + + +async def test_unknown_tool_is_returned_and_tools_are_omitted_when_empty(tmp_path: Path) -> None: + completion: Final = ScriptedCompletion( + ( + model_response(tool_calls=(function_call("missing", "{}"),)), + model_response(content="Unknown tool handled"), + ) + ) + ctx: Final = make_context(tmp_path) + handler: Final = make_handler(completion) + await handler.start(ctx) + + events: Final = await run_turn(handler, ctx, "Call a missing tool") + + assert "tools" not in completion.calls[0] + assert ( + ToolResult( + id="call-1", + output="ValueError: unknown tool 'missing'", + is_error=True, + ) + in events + ) + assert ctx.final_text == "Unknown tool handled" + + +@pytest.mark.parametrize( + "arguments", + ['{"a": 2}', '{"a": 2, "b": 3, "extra": 4}'], +) +async def test_invalid_tool_arguments_return_validation_error( + tmp_path: Path, + arguments: str, +) -> None: + completion: Final = ScriptedCompletion( + ( + model_response(tool_calls=(function_call("add", arguments),)), + model_response(content="Arguments were invalid"), + ) + ) + ctx: Final = make_context(tmp_path, tools=(add,)) + handler: Final = make_handler(completion) + await handler.start(ctx) + + events: Final = await run_turn(handler, ctx, "Add") + result: Final = next(event for event in events if isinstance(event, ToolResult)) + + assert result.is_error + assert result.output.startswith("ValidationError:") + assert completion.calls[1]["messages"][-1]["content"] == result.output + + +async def test_malformed_tool_arguments_return_error_without_calling_tool(tmp_path: Path) -> None: + completion: Final = ScriptedCompletion( + ( + model_response(tool_calls=(function_call("add", "{"),)), + model_response(content="Malformed arguments"), + ) + ) + ctx: Final = make_context(tmp_path, tools=(add,)) + handler: Final = make_handler(completion) + await handler.start(ctx) + + events: Final = await run_turn(handler, ctx, "Add") + result: Final = next(event for event in events if isinstance(event, ToolResult)) + + assert result.is_error + assert result.output.startswith("JSONDecodeError:") + assert completion.calls[1]["messages"][-1]["content"] == result.output + + +async def test_ask_permission_denial_skips_tool_and_returns_reason(tmp_path: Path) -> None: + completion: Final = ScriptedCompletion( + ( + model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)), + model_response(content="Denied"), + ) + ) + ctx: Final = make_context(tmp_path, tools=(add,), permissions="ask") + handler: Final = make_handler(completion) + await handler.start(ctx) + + events: Final = await run_turn(handler, ctx, "Add", allow=False) + + assert any(isinstance(event, Approval) for event in events) + assert ToolResult(id="call-1", output="denied: not approved", is_error=True) in events + assert completion.calls[1]["messages"][-1]["content"] == "denied: not approved" + + +async def test_ask_permission_allow_runs_tool(tmp_path: Path) -> None: + completion: Final = ScriptedCompletion( + ( + model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)), + model_response(content="Allowed"), + ) + ) + ctx: Final = make_context(tmp_path, tools=(add,), permissions="ask") + handler: Final = make_handler(completion) + await handler.start(ctx) + + events: Final = await run_turn(handler, ctx, "Add") + + assert any(isinstance(event, Approval) for event in events) + assert ToolResult(id="call-1", output="5", is_error=False) in events + + +async def test_async_tool_is_awaited(tmp_path: Path) -> None: + async def multiply(a: int, b: int) -> int: + return a * b + + completion: Final = ScriptedCompletion( + ( + model_response(tool_calls=(function_call("multiply", '{"a": 3, "b": 4}'),)), + model_response(content="12"), + ) + ) + ctx: Final = make_context(tmp_path, tools=(multiply,)) + handler: Final = make_handler(completion) + await handler.start(ctx) + + events: Final = await run_turn(handler, ctx, "Multiply") + + assert ToolResult(id="call-1", output="12", is_error=False) in events + + +class Answer(BaseModel): + value: int + + +async def test_structured_output_is_forwarded_and_retained(tmp_path: Path) -> None: + completion: Final = ScriptedCompletion((model_response(content='{"value": 7}'),)) + ctx: Final = make_context(tmp_path, output=Answer) + handler: Final = make_handler(completion) + await handler.start(ctx) + + await run_turn(handler, ctx, "Return a value") + + assert completion.calls[0]["response_format"] is Answer + assert ctx.output_json == '{"value": 7}' + assert ctx.final_text == '{"value": 7}' + + +async def test_gateway_routing_uses_proxy_model_and_tool_loop_tag(tmp_path: Path) -> None: + completion: Final = ScriptedCompletion((model_response(content="done"),)) + gateway: Final = GatewayTarget(api_base="http://gateway", api_key="sk-virtual") + ctx: Final = make_context(tmp_path, gateway=gateway, metadata={"team": "test"}) + handler: Final = make_handler(completion) + await handler.start(ctx) + + await run_turn(handler, ctx, "Hi") + + assert completion.calls[0]["model"] == "litellm_proxy/gpt-4o-mini" + assert completion.calls[0]["api_base"] == "http://gateway" + assert completion.calls[0]["api_key"] == "sk-virtual" + headers: Final = completion.calls[0]["extra_headers"] + assert isinstance(headers, dict) + assert headers["x-litellm-tags"] == "harness,tool_loop" + assert '"team": "test"' in headers["x-litellm-spend-logs-metadata"] + + +async def test_usage_and_cost_accumulate_across_model_calls(tmp_path: Path) -> None: + completion: Final = ScriptedCompletion( + ( + model_response( + tool_calls=(function_call("add", '{"a": 2, "b": 3}'),), + prompt_tokens=10, + completion_tokens=4, + hidden_params={ + "additional_headers": {"llm_provider-x-litellm-response-cost": "0.4"}, + "response_cost": 0.1, + }, + ), + model_response( + content="done", + prompt_tokens=20, + completion_tokens=5, + hidden_params={"response_cost": 0.2}, + ), + ) + ) + ctx: Final = make_context(tmp_path, tools=(add,)) + handler: Final = make_handler(completion) + await handler.start(ctx) + + await run_turn(handler, ctx, "Add") + + assert ctx.calls == 2 + assert ctx.input_tokens == 30 + assert ctx.output_tokens == 9 + assert ctx.cost == pytest.approx(0.6) + + +async def test_history_survives_stop_and_start(tmp_path: Path) -> None: + completion: Final = ScriptedCompletion((model_response(content="first"), model_response(content="second"))) + ctx: Final = make_context(tmp_path, instructions="Keep answers concise") + handler: Final = make_handler(completion) + await handler.start(ctx) + await run_turn(handler, ctx, "first prompt") + await handler.stop(ctx) + await handler.start(ctx) + + await run_turn(handler, ctx, "second prompt") + + assert completion.calls[1]["messages"] == [ + {"role": "system", "content": "Keep answers concise"}, + {"role": "user", "content": "first prompt"}, + {"role": "assistant", "content": "first"}, + {"role": "user", "content": "second prompt"}, + ] + history: Final = await handler.history(ctx) + history[0]["content"] = "changed" + assert (await handler.history(ctx))[0]["content"] == "Keep answers concise" + + +async def test_duplicate_tool_names_are_rejected(tmp_path: Path) -> None: + ctx: Final = make_context(tmp_path, tools=(add, duplicate_add())) + handler: Final = make_handler(ScriptedCompletion(())) + + with pytest.raises(ValueError, match="tool names must be unique"): + await handler.start(ctx) + + +async def test_model_call_limit_raises_harness_turn_error(tmp_path: Path) -> None: + repeating_response: Final = model_response(tool_calls=(function_call("ping", "{}"),)) + completion: Final = ScriptedCompletion((repeating_response,) * 100) + + def ping() -> str: + return "pong" + + ctx: Final = make_context(tmp_path, tools=(ping,)) + handler: Final = make_handler(completion) + await handler.start(ctx) + + with pytest.raises(HarnessTurnError, match="exceeded 100 model calls"): + await run_turn(handler, ctx, "Ping repeatedly") + + +async def test_public_aagent_uses_tool_loop_with_mock_response(tmp_path: Path) -> None: + result: Final = await litellm.aagent( + Harness.TOOL_LOOP, + "Say done", + sandbox=sandbox.local(tmp_path), + model="gpt-4o-mini", + options=ToolLoopOptions(completion_kwargs={"mock_response": "done"}), + ) + + assert result.text == "done" + assert result.stop_reason == "done" diff --git a/tests/unit/harness/test_init.py b/tests/unit/harness/test_init.py index effdc238249..75ec3b60c74 100644 --- a/tests/unit/harness/test_init.py +++ b/tests/unit/harness/test_init.py @@ -35,6 +35,7 @@ PUBLIC_NAMES = [ "CodexOptions", "OpenCodeOptions", "DeepAgentsOptions", + "ToolLoopOptions", "HarnessError", "CapabilityUnsupported", "OptionsMismatch", @@ -89,7 +90,9 @@ def test_adapter_registry_paths_cover_every_harness(): def test_litellm_agent_is_top_level_and_lazy(): code = ( "import sys, litellm; assert 'litellm.harness' not in sys.modules; " - "assert litellm.agent is litellm.harness.agent; assert litellm.Harness.CODEX.value == 'codex'" + "assert litellm.agent is litellm.harness.agent; " + "assert litellm.Harness.CODEX.value == 'codex'; " + "assert litellm.ToolLoopOptions is litellm.harness.ToolLoopOptions" ) out = run_child_interpreter(code, timeout=120) assert out.returncode == 0, out.stderr diff --git a/tests/unit/harness/test_types.py b/tests/unit/harness/test_types.py index 036d6dd625c..e59c4e983fb 100644 --- a/tests/unit/harness/test_types.py +++ b/tests/unit/harness/test_types.py @@ -20,6 +20,7 @@ from litellm.harness.types import ( def test_harness_is_plain_enum(): assert Harness.CODEX.value == "codex" + assert Harness.TOOL_LOOP.value == "tool_loop" assert not isinstance(Harness.CODEX, str) diff --git a/tests/unit/llms/tool_loop/__init__.py b/tests/unit/llms/tool_loop/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/tool_loop/harness/__init__.py b/tests/unit/llms/tool_loop/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/tool_loop/harness/test_transformation.py b/tests/unit/llms/tool_loop/harness/test_transformation.py new file mode 100644 index 00000000000..11bce41bb56 --- /dev/null +++ b/tests/unit/llms/tool_loop/harness/test_transformation.py @@ -0,0 +1,154 @@ +"""Tests for Tool Loop schemas and completion routing.""" + +from __future__ import annotations + +from collections.abc import Callable +from pathlib import Path +from typing import Final, Literal + +import pytest +from pydantic import BaseModel + +from litellm import sandbox +from litellm.harness.context import GatewayTarget, SessionContext +from litellm.harness.options import ToolLoopOptions +from litellm.harness.types import Harness +from litellm.llms.tool_loop.harness.transformation import ( + ToolLoopHarnessConfig, + completion_kwargs, + function_tool, +) + + +def make_context( + tmp_path: Path, + *, + model: str | None = "anthropic/claude", + gateway: GatewayTarget | None = None, + api_key: str | None = None, + api_base: str | None = None, + options: ToolLoopOptions | None = None, + output: type[BaseModel] | None = None, +) -> SessionContext: + return SessionContext( + harness=Harness.TOOL_LOOP, + sandbox=sandbox.local(tmp_path), + session_id="tool-loop-transform", + model=model, + gateway=gateway, + api_key=api_key, + api_base=api_base, + options=options, + output=output, + ) + + +def search( + query: str, + limit: int = 5, + state: Literal["open", "closed"] = "open", +) -> str: + """Search records.""" + return query + + +def test_function_tool_schema_has_required_defaulted_and_literal_fields() -> None: + specification: Final = function_tool(search).spec + schema: Final = specification["function"]["parameters"] + + assert schema["required"] == ["query"] + assert schema["properties"]["query"] == {"title": "Query", "type": "string"} + assert schema["properties"]["limit"] == {"default": 5, "title": "Limit", "type": "integer"} + assert schema["properties"]["state"] == { + "default": "open", + "enum": ["open", "closed"], + "title": "State", + "type": "string", + } + assert specification["function"]["description"] == "Search records." + + +def test_function_schema_rejects_unknown_arguments() -> None: + with pytest.raises(ValueError, match="Extra inputs are not permitted"): + function_tool(search).args_model.model_validate({"query": "owner", "unknown": "value"}) + + +def variadic_positional(*args: int) -> int: + return len(args) + + +def variadic_keyword(**kwargs: int) -> int: + return len(kwargs) + + +@pytest.mark.parametrize("fn", [variadic_positional, variadic_keyword]) +def test_variadic_tools_are_rejected(fn: Callable[..., object]) -> None: + with pytest.raises(ValueError, match="variadic parameters"): + function_tool(fn) + + +def test_sdk_routing_overrides_completion_kwargs(tmp_path: Path) -> None: + ctx: Final = make_context( + tmp_path, + api_key="provided-key", + api_base="https://provider", + options=ToolLoopOptions( + completion_kwargs={ + "model": "wrong-model", + "api_key": "wrong-key", + "api_base": "https://wrong", + "temperature": 0.2, + } + ), + ) + + kwargs: Final = completion_kwargs(ctx) + + assert kwargs == { + "model": "anthropic/claude", + "api_key": "provided-key", + "api_base": "https://provider", + "temperature": 0.2, + } + + +def test_gateway_routing_and_response_format_override_options(tmp_path: Path) -> None: + class OutputModel(BaseModel): + pass + + gateway: Final = GatewayTarget(api_base="https://gateway", api_key="virtual-key") + ctx: Final = make_context( + tmp_path, + gateway=gateway, + options=ToolLoopOptions( + completion_kwargs={ + "model": "wrong-model", + "api_key": "wrong-key", + "api_base": "https://wrong", + "response_format": "wrong-format", + } + ), + output=OutputModel, + ) + + kwargs: Final = completion_kwargs(ctx) + + assert kwargs == { + "model": "litellm_proxy/anthropic/claude", + "api_base": "https://gateway", + "api_key": "virtual-key", + "extra_headers": {"x-litellm-tags": "harness,tool_loop"}, + "response_format": OutputModel, + } + + +def test_configuration_requires_model_and_declares_capabilities(tmp_path: Path) -> None: + config: Final = ToolLoopHarnessConfig() + assert config.uses_model_endpoint is False + assert config.capabilities.structured_output + assert config.capabilities.tool_approval + assert config.capabilities.history + assert config.capabilities.custom_tools + assert config.capabilities.permission_modes == frozenset({"ask", "full"}) + with pytest.raises(ValueError, match=r"Harness\.TOOL_LOOP needs model="): + config.validate_environment(make_context(tmp_path, model=None)) From 8b1990b4bcc61a98da08bb71e7519d6b53c8572a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 17:38:38 +0000 Subject: [PATCH 03/82] feat(decisions): add unified /v1/decisions endpoint for Jev-compatible providers (#44236) * feat(decisions): add unified /v1/decisions endpoint for Jev-compatible providers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(decisions): register typesafe as a provider so Jev deployments load Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(decisions): move provider endpoints under llms and validate proxy bodies Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(decisions): add Cloudflare Clef and Strands Decider backends Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(decisions): register decisions routes for managed agents and gateway Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(decisions): use raw regex for cloudflare missing account match Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(decisions): avoid cast in Cloudflare response unwrapping Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(decisions): default model, evaluation health probe, short Cloudflare names The proxy validates only state and questions, so a request without a model falls through to the configured default model like every other route. Health checks probe evaluation-mode deployments through the Decisions API instead of failing with an unsupported mode, and cloudflare/clef and cloudflare/clef-flash get cost-map rows so the short names resolve a mode and a price. The registry no longer claims typed decisions for a provider with no backend. * fix(decisions): let health_check_params override the evaluation probe Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): audit the decisions endpoint across providers, limits, health and chaos Adds the /v1/decisions audit cells: one wire contract per provider (path, key, body and cost-map billing), the gateway-only fields and tags, the sad paths (invalid bodies, unknown model, key checks, api_base in the body, upstream 401/429/500, a 200 without answers, an unreachable upstream), the two evaluation-mode health probes, and three chaos cells (a mixed-failure burst over both routes, a worker SIGKILL mid-burst, an upstream outage and restart on the same port). The PR's cost case read the upstream observations through the gateway, which answers 404 for that path; it now reads them from the upstream URL. The owned proxy harness takes extra CLI arguments, and its graceful stop waits as long as a worker boot may take, since a worker still starting honors SIGTERM only once it is up and the 30 second wait forced a cleanup under load. * fix(decisions): send env API keys to a configured api_base Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(decisions): add zero-cost evaluation cost-map entry for Strands Decider Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(decisions): register the routes through the lazy feature registry The Decisions router was included at import, ahead of the config and DB pass-through endpoints, so a pass-through configured at /v1/decisions was skipped and answered 400 as an unknown Decisions provider. The routes now register through LAZY_FEATURES, which splices them in after every eager route, so a pass-through at /v1/decisions keeps its route while /decisions still serves natively. The lazy OpenAPI snapshot carries the two paths so the schema shows them before the first call. The audit cells add the env-key egress to a configured api_base, the client api_base opt-in shared with chat, the pass-through precedence on an owned proxy, and the Strands evaluation health check resolved from the cost map. The integration config exports the Perplexity env key the first cell needs. * fix(decisions): keep the Cloudflare api_base message in its transformation and read the audit upstream once per cell * fix(proxy): let a config pass-through beat a lazily registered route in eager mode With LITELLM_DISABLE_LAZY_ROUTES set the decisions routes are registered at startup, so SafeRouteAdder treated a config pass-through at exactly /v1/decisions as already registered and dropped it. In lazy mode a pass-through created through the API after the first native call was skipped the same way. Routes a lazy feature owns no longer count as registered, and a route added at one of their paths is placed ahead of them, the precedence lazy mode gives a config pass-through when the feature has not loaded yet. --------- Co-authored-by: mateo Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .github/workflows/test-unit.yml | 3 +- README.md | 2 + gateway/routes/allowlist.py | 4 +- litellm/__init__.py | 1 + litellm/cost_calculator.py | 14 +- litellm/decisions/__init__.py | 3 + litellm/decisions/main.py | 299 +++++++++ .../health_check_helpers.py | 12 +- .../litellm_core_utils/health_check_utils.py | 6 + litellm/litellm_core_utils/litellm_logging.py | 3 + litellm/llms/base_llm/decisions/__init__.py | 3 + .../llms/base_llm/decisions/transformation.py | 52 ++ .../cloudflare/decisions/transformation.py | 58 ++ .../openrouter/decisions/transformation.py | 10 + .../perplexity/decisions/transformation.py | 10 + .../decisions/transformation.py | 11 + .../llms/typesafe/decisions/transformation.py | 10 + ...odel_prices_and_context_window_backup.json | 50 ++ litellm/proxy/_lazy_features.py | 27 +- litellm/proxy/_lazy_openapi_snapshot.json | 55 ++ litellm/proxy/_types.py | 2 + .../auth/managed_authorization.py | 2 + litellm/proxy/common_request_processing.py | 2 + litellm/proxy/decisions_endpoints/__init__.py | 1 + .../proxy/decisions_endpoints/endpoints.py | 102 +++ .../pass_through_endpoints.py | 50 +- litellm/proxy/route_llm_request.py | 2 + litellm/router.py | 9 + litellm/types/decisions.py | 123 ++++ litellm/types/utils.py | 8 + litellm/utils.py | 3 + model_prices_and_context_window.json | 50 ++ provider_endpoints_support.json | 14 + tests/integration/_support/process.py | 8 +- .../cost_calculation/cost_tracking_case.py | 1 + .../cost_calculation/cost_tracking_cases.json | 48 ++ .../cost_calculation/test_cost_tracking.py | 25 + .../management/test_model_health_check.py | 117 +++- .../providers/test_decisions_chaos.py | 262 ++++++++ .../providers/test_decisions_wire.py | 456 ++++++++++++++ tests/integration/proxy_config.yaml | 2 + tests/unit/decisions/__init__.py | 0 tests/unit/decisions/test_main.py | 592 ++++++++++++++++++ .../test_health_check_helpers.py | 99 +++ .../llms/azure/test_azure_common_utils.py | 1 + .../proxy/decisions_endpoints/__init__.py | 0 .../decisions_endpoints/test_endpoints.py | 326 ++++++++++ .../test_pass_through_endpoints.py | 35 +- .../public_endpoints/test_public_endpoints.py | 2 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 76 ++- 50 files changed, 3011 insertions(+), 40 deletions(-) create mode 100644 litellm/decisions/__init__.py create mode 100644 litellm/decisions/main.py create mode 100644 litellm/llms/base_llm/decisions/__init__.py create mode 100644 litellm/llms/base_llm/decisions/transformation.py create mode 100644 litellm/llms/cloudflare/decisions/transformation.py create mode 100644 litellm/llms/openrouter/decisions/transformation.py create mode 100644 litellm/llms/perplexity/decisions/transformation.py create mode 100644 litellm/llms/strands_decider/decisions/transformation.py create mode 100644 litellm/llms/typesafe/decisions/transformation.py create mode 100644 litellm/proxy/decisions_endpoints/__init__.py create mode 100644 litellm/proxy/decisions_endpoints/endpoints.py create mode 100644 litellm/types/decisions.py create mode 100644 tests/integration/providers/test_decisions_chaos.py create mode 100644 tests/integration/providers/test_decisions_wire.py create mode 100644 tests/unit/decisions/__init__.py create mode 100644 tests/unit/decisions/test_main.py create mode 100644 tests/unit/proxy/decisions_endpoints/__init__.py create mode 100644 tests/unit/proxy/decisions_endpoints/test_endpoints.py diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 06c2990a4a4..b769c8f6a3d 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -61,7 +61,7 @@ jobs: - shard: core-utils artifact-name: core-utils - test-path: "" + test-path: tests/unit/decisions unit-flag: core-utils workers: 2 reruns: 1 @@ -141,6 +141,7 @@ jobs: artifact-name: proxy-endpoints test-path: >- tests/unit/proxy/analytics_endpoints + tests/unit/proxy/decisions_endpoints tests/unit/proxy/management_endpoints tests/unit/proxy/list_api tests/unit/proxy/memory diff --git a/README.md b/README.md index 4004e6474ee..7ffc44854bb 100644 --- a/README.md +++ b/README.md @@ -390,11 +390,13 @@ Set `LITELLM_PROXY_API_BASE` and `LITELLM_PROXY_API_KEY` and every model call th | [Sail (`sail`)](https://docs.litellm.ai/docs/providers/sail) | ✅ | ✅ | ✅ | | | | | | | | | [Sambanova (`sambanova`)](https://docs.litellm.ai/docs/providers/sambanova) | ✅ | ✅ | ✅ | | | | | | | | | [Snowflake (`snowflake`)](https://docs.litellm.ai/docs/providers/snowflake) | ✅ | ✅ | ✅ | | | | | | | | +| [Strands Decider (`strands_decider`)](https://docs.litellm.ai/docs/providers) | | | | | | | | | | | | [Text Completion Codestral (`text-completion-codestral`)](https://docs.litellm.ai/docs/providers/codestral) | ✅ | ✅ | ✅ | | | | | | | | | [Text Completion OpenAI (`text-completion-openai`)](https://docs.litellm.ai/docs/providers/text_completion_openai) | ✅ | ✅ | ✅ | | | ✅ | ✅ | ✅ | ✅ | | | [Together AI (`together_ai`)](https://docs.litellm.ai/docs/providers/togetherai) | ✅ | ✅ | ✅ | | | | | | | | | [Topaz (`topaz`)](https://docs.litellm.ai/docs/providers/topaz) | ✅ | ✅ | ✅ | | | | | | | | | [Triton (`triton`)](https://docs.litellm.ai/docs/providers/triton-inference-server) | ✅ | ✅ | ✅ | | | | | | | | +| [Typesafe Decisions API (`typesafe`)](https://docs.litellm.ai/docs/providers) | | | | | | | | | | | | [V0 (`v0`)](https://docs.litellm.ai/docs/providers/v0) | ✅ | ✅ | ✅ | | | | | | | | | [Vercel AI Gateway (`vercel_ai_gateway`)](https://docs.litellm.ai/docs/providers/vercel_ai_gateway) | ✅ | ✅ | ✅ | | | | | | | | | [VLLM (`vllm`)](https://docs.litellm.ai/docs/providers/vllm) | ✅ | ✅ | ✅ | | | | | | | | diff --git a/gateway/routes/allowlist.py b/gateway/routes/allowlist.py index 6e91f5486d0..fc11c059c85 100644 --- a/gateway/routes/allowlist.py +++ b/gateway/routes/allowlist.py @@ -1,7 +1,7 @@ """Path allowlist for the gateway component. The gateway exposes the LLM data-plane surface: chat/completions, embeddings, -audio, batches, files, fine-tuning, rerank, ocr, rag, video, search, image, +audio, batches, files, fine-tuning, rerank, decisions, ocr, rag, video, search, image, responses, vector stores, passthrough providers, realtime websockets, MCP tool-call endpoints, and operational endpoints (/health, /metrics, and the /debug/memory/summary read of the serving worker's RSS). @@ -60,6 +60,8 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = ( "/v1/rerank", "/v2/rerank", "/rerank", + "/v1/decisions", + "/decisions", "/v1/ocr", "/ocr", "/v1/rag/", diff --git a/litellm/__init__.py b/litellm/__init__.py index 30a9a80e9b2..b6428b51bfb 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1473,6 +1473,7 @@ from .embeddings.dispatch import * from .rust_bridge import rust from .rag.main import * from .sandbox.main import * +from .decisions.main import * from .search.main import * from .realtime_api.main import ( _arealtime, diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 41a7ef1ab64..e358636c105 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -99,6 +99,7 @@ from litellm.llms.vertex_ai.cost_calculator import cost_router as google_cost_ro from litellm.llms.xai.cost_calculator import cost_per_token as xai_cost_per_token from litellm.responses.utils import ResponseAPILoggingUtils from litellm.types.agents import LiteLLMSendMessageResponse +from litellm.types.decisions import DecisionsResponse, DecisionsUsage from litellm.types.llms.base import CachedTokensDetails from litellm.types.llms.openai import ( HttpxBinaryResponseContent, @@ -1058,6 +1059,7 @@ def _is_known_usage_objects(usage_obj): return ( isinstance(usage_obj, litellm.Usage) or isinstance(usage_obj, ResponseAPIUsage) + or isinstance(usage_obj, DecisionsUsage) or TranscriptionUsageObjectTransformation.is_transcription_usage_object(usage_obj) ) @@ -1466,7 +1468,12 @@ def completion_cost( "usage", litellm.Usage(**_usage_for_dump.model_dump()), ) - if usage_obj is None: + if isinstance(usage_obj, DecisionsUsage): + _usage = { + "prompt_tokens": usage_obj.input_tokens, + "completion_tokens": usage_obj.output_tokens, + } + elif usage_obj is None: _usage = {} elif isinstance(usage_obj, BaseModel): _usage = cast(BaseModel, usage_obj).model_dump() @@ -1957,7 +1964,8 @@ def response_cost_calculator( | LiteLLMRealtimeStreamLoggingObject | OpenAIModerationResponse | Response - | SearchResponse, + | SearchResponse + | DecisionsResponse, model: str, custom_llm_provider: str | None, call_type: Literal[ @@ -1979,6 +1987,8 @@ def response_cost_calculator( "arerank", "search", "asearch", + "decisions", + "adecisions", ], optional_params: dict, cache_hit: bool | None = None, diff --git a/litellm/decisions/__init__.py b/litellm/decisions/__init__.py new file mode 100644 index 00000000000..40bddd200a9 --- /dev/null +++ b/litellm/decisions/__init__.py @@ -0,0 +1,3 @@ +from litellm.decisions.main import adecisions, decisions + +__all__ = ["adecisions", "decisions"] diff --git a/litellm/decisions/main.py b/litellm/decisions/main.py new file mode 100644 index 00000000000..da037f1d8cb --- /dev/null +++ b/litellm/decisions/main.py @@ -0,0 +1,299 @@ +from collections.abc import Mapping +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import Final + +import httpx +from pydantic import TypeAdapter, ValidationError + +import litellm +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.decisions.transformation import DecisionsProviderConfig +from litellm.llms.cloudflare.decisions.transformation import CLOUDFLARE_DECISIONS_ENDPOINT +from litellm.llms.custom_httpx.http_handler import _get_httpx_client, get_async_httpx_client +from litellm.llms.openrouter.decisions.transformation import OPENROUTER_DECISIONS_ENDPOINT +from litellm.llms.perplexity.decisions.transformation import PERPLEXITY_DECISIONS_ENDPOINT +from litellm.llms.strands_decider.decisions.transformation import STRANDS_DECIDER_DECISIONS_ENDPOINT +from litellm.llms.typesafe.decisions.transformation import TYPESAFE_DECISIONS_ENDPOINT +from litellm.secret_managers.main import get_secret_str +from litellm.types.decisions import ( + DecisionQuestion, + DecisionsJSON, + DecisionsRequest, + DecisionsResponse, +) +from litellm.utils import client + +DECISIONS_ENDPOINTS: Final[Mapping[str, DecisionsProviderConfig]] = MappingProxyType( + { + "perplexity": PERPLEXITY_DECISIONS_ENDPOINT, + "typesafe": TYPESAFE_DECISIONS_ENDPOINT, + "openrouter": OPENROUTER_DECISIONS_ENDPOINT, + "cloudflare": CLOUDFLARE_DECISIONS_ENDPOINT, + "strands_decider": STRANDS_DECIDER_DECISIONS_ENDPOINT, + } +) + +_DECISIONS_REQUEST_ADAPTER: Final[TypeAdapter[DecisionsRequest]] = TypeAdapter(DecisionsRequest) +_DECISIONS_PAYLOAD_ADAPTER: Final[TypeAdapter[object]] = TypeAdapter(object) +_DECISIONS_RESPONSE_ADAPTER: Final[TypeAdapter[DecisionsResponse]] = TypeAdapter(DecisionsResponse) + + +@dataclass(frozen=True, slots=True, repr=False) +class _PreparedDecisionsRequest: + config: DecisionsProviderConfig + provider: str + upstream_model: str + url: str + api_key: str | None = field(repr=False) + headers: Mapping[str, str] = field(repr=False) + body: Mapping[str, object] = field(repr=False) + + +def _resolve_provider_model(model: str, custom_llm_provider: str | None) -> tuple[str, str]: + provider: Final = model.partition("/")[0] if custom_llm_provider is None else custom_llm_provider + if provider not in DECISIONS_ENDPOINTS: + supported: Final = ", ".join(DECISIONS_ENDPOINTS) + raise litellm.BadRequestError( + message=f"Unknown Decisions provider '{provider}'. Supported providers: {supported}", + model=model, + llm_provider=provider, + ) + upstream_model: Final = model.removeprefix(f"{provider}/") if model.startswith(f"{provider}/") else model + if not upstream_model: + raise litellm.BadRequestError( + message="A model name is required for the Decisions API", + model=model, + llm_provider=provider, + ) + return provider, upstream_model + + +def _resolve_api_key( + *, + provider: str, + model: str, + endpoint: DecisionsProviderConfig, + api_key: str | None, +) -> str | None: + if api_key is not None: + return api_key + + server_api_key: Final = next( + (key for key in (get_secret_str(name) for name in endpoint.api_key_env) if key), + None, + ) + if server_api_key is None: + if not endpoint.api_key_required: + return None + raise litellm.AuthenticationError( + message=f"Missing API key for Decisions provider '{provider}'", + model=model, + llm_provider=provider, + ) + + return server_api_key + + +def _prepare_request( + *, + model: str, + state: DecisionsJSON, + questions: Mapping[str, DecisionQuestion | Mapping[str, object]], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, str] | None, +) -> _PreparedDecisionsRequest: + provider, upstream_model = _resolve_provider_model(model, custom_llm_provider) + try: + validated_request: Final = _DECISIONS_REQUEST_ADAPTER.validate_python( + {"model": model, "state": state, "questions": questions} + ) + except ValidationError as error: + raise litellm.BadRequestError( + message=f"Invalid Decisions request: {error}", + model=model, + llm_provider=provider, + ) from error + + endpoint: Final = DECISIONS_ENDPOINTS[provider] + env_api_base: Final = get_secret_str(endpoint.api_base_env) + default_api_base: Final = endpoint.default_api_base() + resolved_api_base: Final = api_base or env_api_base or default_api_base + if resolved_api_base is None: + raise litellm.BadRequestError( + message=endpoint.missing_api_base_message(provider), + model=model, + llm_provider=provider, + ) + + resolved_api_key: Final = _resolve_api_key( + provider=provider, + model=model, + endpoint=endpoint, + api_key=api_key, + ) + + canonical_model: Final = endpoint.canonical_model(upstream_model) + outbound_headers: Final = MappingProxyType( + { + **{ + name: value + for name, value in (extra_headers or {}).items() + if name.lower() not in {"authorization", "content-type"} + }, + **({"Authorization": f"Bearer {resolved_api_key}"} if resolved_api_key is not None else {}), + "Content-Type": "application/json", + } + ) + body: Final = MappingProxyType( + { + "model": endpoint.request_model(canonical_model), + "state": validated_request.state, + "questions": { + name: question.model_dump(mode="json", exclude_none=True) + for name, question in validated_request.questions.items() + }, + } + ) + return _PreparedDecisionsRequest( + config=endpoint, + provider=provider, + upstream_model=canonical_model, + url=endpoint.endpoint_url(resolved_api_base, canonical_model), + api_key=resolved_api_key, + headers=outbound_headers, + body=body, + ) + + +def _log_request( + prepared: _PreparedDecisionsRequest, + kwargs: Mapping[str, object], +) -> LiteLLMLoggingObj | None: + logging_obj: Final = kwargs.get("litellm_logging_obj") + if not isinstance(logging_obj, LiteLLMLoggingObj): + return None + logging_obj.update_from_kwargs( + kwargs=dict(kwargs), + model=prepared.upstream_model, + litellm_params={ + "litellm_call_id": kwargs.get("litellm_call_id"), + "api_base": prepared.url, + }, + custom_llm_provider=prepared.provider, + ) + request_body: Final = dict(prepared.body) + request_headers: Final = dict(prepared.headers) + logging_obj.pre_call( + input=request_body, + api_key=prepared.api_key, + model=prepared.upstream_model, + additional_args={ + "api_base": prepared.url, + "complete_input_dict": request_body, + "headers": request_headers, + }, + ) + return logging_obj + + +def _parse_response( + response: httpx.Response, + prepared: _PreparedDecisionsRequest, +) -> DecisionsResponse: + response.raise_for_status() + payload: Final[object] = _DECISIONS_PAYLOAD_ADAPTER.validate_json(response.content) + result: Final = _DECISIONS_RESPONSE_ADAPTER.validate_python(prepared.config.unwrap_response(payload)) + result._hidden_params.update( + { + "model": f"{prepared.provider}/{prepared.upstream_model}", + "custom_llm_provider": prepared.provider, + "provider_response_model": f"{prepared.provider}/{prepared.upstream_model}", + } + ) + return result + + +def _map_upstream_exception(error: Exception, prepared: _PreparedDecisionsRequest) -> Exception: + return litellm.exception_type( + model=f"{prepared.provider}/{prepared.upstream_model}", + custom_llm_provider=prepared.provider, + original_exception=error, + ) + + +@client +async def adecisions( + model: str, + state: DecisionsJSON, + questions: Mapping[str, DecisionQuestion | Mapping[str, object]], + api_key: str | None = None, + api_base: str | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, + extra_headers: Mapping[str, str] | None = None, + **kwargs: object, +) -> DecisionsResponse: + prepared: Final = _prepare_request( + model=model, + state=state, + questions=questions, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + ) + logging_obj: Final = _log_request(prepared, kwargs) + try: + handler: Final = get_async_httpx_client(llm_provider=prepared.provider) + response: Final = await handler.post( + prepared.url, + json=dict(prepared.body), + headers=dict(prepared.headers), + timeout=timeout, + logging_obj=logging_obj, + ) + return _parse_response(response=response, prepared=prepared) + except Exception as error: + raise _map_upstream_exception(error, prepared) from error + + +@client +def decisions( + model: str, + state: DecisionsJSON, + questions: Mapping[str, DecisionQuestion | Mapping[str, object]], + api_key: str | None = None, + api_base: str | None = None, + timeout: float | httpx.Timeout | None = None, + custom_llm_provider: str | None = None, + extra_headers: Mapping[str, str] | None = None, + **kwargs: object, +) -> DecisionsResponse: + prepared: Final = _prepare_request( + model=model, + state=state, + questions=questions, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + ) + logging_obj: Final = _log_request(prepared, kwargs) + try: + handler: Final = _get_httpx_client() + response: Final = handler.post( + prepared.url, + json=dict(prepared.body), + headers=dict(prepared.headers), + timeout=timeout, + logging_obj=logging_obj, + ) + return _parse_response(response=response, prepared=prepared) + except Exception as error: + raise _map_upstream_exception(error, prepared) from error + + +__all__ = ["DECISIONS_ENDPOINTS", "adecisions", "decisions"] diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index a0f027cd58f..41965404351 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -168,6 +168,7 @@ class HealthCheckHelpers: "batch", "responses", "ocr", + "evaluation", ], Callable, ]: @@ -190,7 +191,7 @@ class HealthCheckHelpers: from litellm.litellm_core_utils.audio_utils.utils import ( get_audio_file_for_health_check, ) - from litellm.litellm_core_utils.health_check_utils import _filter_model_params + from litellm.litellm_core_utils.health_check_utils import DECISIONS_CALL_PARAMS, _filter_model_params from litellm.realtime_api.main import _realtime_health_check return { @@ -257,4 +258,13 @@ class HealthCheckHelpers: **_filter_model_params(model_params=model_params), document=_ocr_health_check_document(model=model, custom_llm_provider=custom_llm_provider), ), + "evaluation": lambda: litellm.adecisions( + **DECISIONS_CALL_PARAMS.validate_python( + { + "state": prompt or "health check", + "questions": {"reachable": {"type": "noul", "instructions": "Is the service reachable?"}}, + **_filter_model_params(model_params=model_params), + } + ) + ), } diff --git a/litellm/litellm_core_utils/health_check_utils.py b/litellm/litellm_core_utils/health_check_utils.py index ae56ae8f899..7fe2d830f1e 100644 --- a/litellm/litellm_core_utils/health_check_utils.py +++ b/litellm/litellm_core_utils/health_check_utils.py @@ -4,6 +4,12 @@ Utils used for litellm.ahealth_check() from typing import Final +from pydantic import TypeAdapter + +from litellm.types.decisions import DecisionsCallParams + +DECISIONS_CALL_PARAMS: Final[TypeAdapter[DecisionsCallParams]] = TypeAdapter(DecisionsCallParams) + def _filter_model_params(model_params: dict) -> dict: """Remove 'messages' param from model params.""" diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 26c02bb0243..5a3a17f338c 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -118,6 +118,7 @@ from litellm.llms.base_llm.search.transformation import SearchResponse from litellm.responses.utils import ResponseAPILoggingUtils from litellm.types.agents import LiteLLMSendMessageResponse from litellm.types.containers.main import ContainerObject +from litellm.types.decisions import DecisionsResponse from litellm.types.integrations.s3_v2 import S3PartitionGranularity from litellm.types.interactions import ( InteractionsAPIResponse, @@ -1815,6 +1816,7 @@ class Logging(LiteLLMLoggingBaseClass): LiteLLMRealtimeStreamLoggingObject, OpenAIModerationResponse, "SearchResponse", + DecisionsResponse, dict, list, ], @@ -2600,6 +2602,7 @@ class Logging(LiteLLMLoggingBaseClass): or isinstance(logging_result, OpenAIModerationResponse) or isinstance(logging_result, OCRResponse) # OCR or isinstance(logging_result, SearchResponse) # Search API + or isinstance(logging_result, DecisionsResponse) or ( isinstance(logging_result, InteractionsAPIResponse) and logging_result.usage is not None diff --git a/litellm/llms/base_llm/decisions/__init__.py b/litellm/llms/base_llm/decisions/__init__.py new file mode 100644 index 00000000000..c18ac9b00f2 --- /dev/null +++ b/litellm/llms/base_llm/decisions/__init__.py @@ -0,0 +1,3 @@ +from .transformation import DecisionsProviderConfig, JevCompatibleDecisionsEndpoint + +__all__ = ["DecisionsProviderConfig", "JevCompatibleDecisionsEndpoint"] diff --git a/litellm/llms/base_llm/decisions/transformation.py b/litellm/llms/base_llm/decisions/transformation.py new file mode 100644 index 00000000000..d4fcea24793 --- /dev/null +++ b/litellm/llms/base_llm/decisions/transformation.py @@ -0,0 +1,52 @@ +from dataclasses import dataclass +from typing import Protocol + + +@dataclass(frozen=True, slots=True) +class JevCompatibleDecisionsEndpoint: + default_api_base_value: str | None + path: str + api_key_env: tuple[str, ...] + api_base_env: str + api_key_required: bool = True + + def default_api_base(self) -> str | None: + return self.default_api_base_value + + def missing_api_base_message(self, provider: str) -> str: + return f"api_base is required for Decisions provider '{provider}'" + + def canonical_model(self, model: str) -> str: + return model + + def request_model(self, model: str) -> str: + return model + + def endpoint_url(self, api_base: str, model: str) -> str: + return f"{api_base.rstrip('/').removesuffix('/v1')}{self.path}" + + def unwrap_response(self, payload: object) -> object: + return payload + + +class DecisionsProviderConfig(Protocol): + @property + def api_key_env(self) -> tuple[str, ...]: ... + + @property + def api_base_env(self) -> str: ... + + @property + def api_key_required(self) -> bool: ... + + def default_api_base(self) -> str | None: ... + + def missing_api_base_message(self, provider: str) -> str: ... + + def canonical_model(self, model: str) -> str: ... + + def request_model(self, model: str) -> str: ... + + def endpoint_url(self, api_base: str, model: str) -> str: ... + + def unwrap_response(self, payload: object) -> object: ... diff --git a/litellm/llms/cloudflare/decisions/transformation.py b/litellm/llms/cloudflare/decisions/transformation.py new file mode 100644 index 00000000000..6e8b2999778 --- /dev/null +++ b/litellm/llms/cloudflare/decisions/transformation.py @@ -0,0 +1,58 @@ +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Final + +from pydantic import TypeAdapter + +from litellm.secret_managers.main import ( + get_secret_str, + normalize_nonempty_secret_str, +) + +_RESPONSE_MAPPING_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object]) + + +@dataclass(frozen=True, slots=True) +class CloudflareDecisionsEndpoint: + api_key_env: tuple[str, ...] = ("CLOUDFLARE_API_KEY",) + api_base_env: str = "CLOUDFLARE_API_BASE" + api_key_required: bool = True + + def default_api_base(self) -> str | None: + account_id: Final = normalize_nonempty_secret_str(get_secret_str("CLOUDFLARE_ACCOUNT_ID")) + if account_id is None: + return None + return f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run" + + def missing_api_base_message(self, provider: str) -> str: + return "Missing CLOUDFLARE_ACCOUNT_ID - set CLOUDFLARE_ACCOUNT_ID or pass api_base explicitly" + + def canonical_model(self, model: str) -> str: + if model.startswith("@cf/"): + return model + return f"@cf/cloudflare/{model}" + + def request_model(self, model: str) -> str: + return model.rsplit("/", maxsplit=1)[-1] + + def endpoint_url(self, api_base: str, model: str) -> str: + normalized_api_base: Final = api_base.rstrip("/") + if normalized_api_base.endswith("/ai/v1"): + return f"{normalized_api_base.removesuffix('/ai/v1')}/ai/run/{model}" + if normalized_api_base.endswith("/ai/run"): + return f"{normalized_api_base}/{model}" + return f"{normalized_api_base}/ai/run/{model}" + + def unwrap_response(self, payload: object) -> object: + if not isinstance(payload, Mapping): + return payload + response_mapping: Final = _RESPONSE_MAPPING_ADAPTER.validate_python(payload) + if "answers" in response_mapping: + return payload + result: Final = response_mapping.get("result") + if isinstance(result, Mapping): + return result + return payload + + +CLOUDFLARE_DECISIONS_ENDPOINT: Final[CloudflareDecisionsEndpoint] = CloudflareDecisionsEndpoint() diff --git a/litellm/llms/openrouter/decisions/transformation.py b/litellm/llms/openrouter/decisions/transformation.py new file mode 100644 index 00000000000..7a7466b1239 --- /dev/null +++ b/litellm/llms/openrouter/decisions/transformation.py @@ -0,0 +1,10 @@ +from typing import Final + +from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint + +OPENROUTER_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint( + default_api_base_value="https://openrouter.ai/api", + path="/alpha/decisions", + api_key_env=("OPENROUTER_API_KEY",), + api_base_env="OPENROUTER_API_BASE", +) diff --git a/litellm/llms/perplexity/decisions/transformation.py b/litellm/llms/perplexity/decisions/transformation.py new file mode 100644 index 00000000000..69a4753f4a3 --- /dev/null +++ b/litellm/llms/perplexity/decisions/transformation.py @@ -0,0 +1,10 @@ +from typing import Final + +from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint + +PERPLEXITY_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint( + default_api_base_value="https://api.perplexity.ai", + path="/v1/decisions", + api_key_env=("PERPLEXITYAI_API_KEY", "PERPLEXITY_API_KEY"), + api_base_env="PERPLEXITY_API_BASE", +) diff --git a/litellm/llms/strands_decider/decisions/transformation.py b/litellm/llms/strands_decider/decisions/transformation.py new file mode 100644 index 00000000000..265afb2b148 --- /dev/null +++ b/litellm/llms/strands_decider/decisions/transformation.py @@ -0,0 +1,11 @@ +from typing import Final + +from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint + +STRANDS_DECIDER_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint( + default_api_base_value=None, + path="/v1/systemone", + api_key_env=("STRANDS_DECIDER_API_KEY",), + api_base_env="STRANDS_DECIDER_API_BASE", + api_key_required=False, +) diff --git a/litellm/llms/typesafe/decisions/transformation.py b/litellm/llms/typesafe/decisions/transformation.py new file mode 100644 index 00000000000..17fb24ce443 --- /dev/null +++ b/litellm/llms/typesafe/decisions/transformation.py @@ -0,0 +1,10 @@ +from typing import Final + +from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint + +TYPESAFE_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint( + default_api_base_value="https://api.typesafe.ai", + path="/v1/systemone", + api_key_env=("TYPESAFE_API_KEY",), + api_base_env="TYPESAFE_API_BASE", +) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index ea383ef4c11..3d2acf4e9d9 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -15612,6 +15612,38 @@ "prompt_cache_min_tokens": 1024, "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, + "cloudflare/clef": { + "input_cost_per_token": 2.4e-07, + "litellm_provider": "cloudflare", + "max_input_tokens": 65536, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://developers.cloudflare.com/workers-ai/platform/pricing/" + }, + "cloudflare/clef-flash": { + "input_cost_per_token": 9e-08, + "litellm_provider": "cloudflare", + "max_input_tokens": 65536, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://developers.cloudflare.com/workers-ai/platform/pricing/" + }, + "cloudflare/@cf/cloudflare/clef": { + "input_cost_per_token": 2.4e-07, + "litellm_provider": "cloudflare", + "max_input_tokens": 65536, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://developers.cloudflare.com/workers-ai/platform/pricing/" + }, + "cloudflare/@cf/cloudflare/clef-flash": { + "input_cost_per_token": 9e-08, + "litellm_provider": "cloudflare", + "max_input_tokens": 65536, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://developers.cloudflare.com/workers-ai/platform/pricing/" + }, "cloudflare/@cf/meta/llama-2-7b-chat-fp16": { "input_cost_per_token": 1.923e-06, "litellm_provider": "cloudflare", @@ -44421,6 +44453,14 @@ "mode": "chat", "output_cost_per_token": 2.8e-07 }, + "perplexity/pplx-decider-v1-27b": { + "input_cost_per_token": 4e-08, + "litellm_provider": "perplexity", + "max_input_tokens": 262144, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://docs.perplexity.ai/api-reference/decisions-post" + }, "perplexity/sonar": { "input_cost_per_token": 1e-06, "litellm_provider": "perplexity", @@ -72800,6 +72840,16 @@ "notes": "Self-hosted decision model; infrastructure costs are paid separately" } }, + "strands_decider/strands-decider-2B-hobson-v19": { + "input_cost_per_token": 0.0, + "litellm_provider": "strands_decider", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://huggingface.co/StrandsAgents/strands-decider-2B-hobson-v19", + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, "typesafe/jev-1.13.0": { "input_cost_per_token": 4.2e-08, "litellm_provider": "typesafe", diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 54d757d75aa..cb6d47ca4c8 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -16,6 +16,7 @@ from collections.abc import Set as AbstractSet from contextlib import asynccontextmanager from dataclasses import dataclass, field from functools import partial +from itertools import chain from types import MappingProxyType from typing import TYPE_CHECKING, Final @@ -268,6 +269,11 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( module_path="litellm.proxy.openai_evals_endpoints.endpoints", path_prefixes=("/v1/evals", "/evals"), ), + LazyFeature( + name="decisions", + module_path="litellm.proxy.decisions_endpoints.endpoints", + path_prefixes=("/v1/decisions", "/decisions"), + ), LazyFeature( name="claude_code_marketplace", module_path="litellm.proxy.anthropic_endpoints.claude_code_endpoints", @@ -378,6 +384,10 @@ def _lazy_slots(app: "FastAPI") -> Mapping[str, BaseRoute | None]: return app.state.lazy_slots if hasattr(app.state, "lazy_slots") else MappingProxyType({}) +def _lazy_routes(app: "FastAPI") -> Mapping[str, tuple[BaseRoute, ...]]: + return app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({}) + + def reserve_lazy_slot(app: "FastAPI", name: str, features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> None: """Record the route the feature's router used to be included after, so its routes are spliced back in there once it loads and keep the same precedence. Anchoring on @@ -474,11 +484,8 @@ def _lazy_lock(app: "FastAPI", module_path: str) -> asyncio.Lock: def _register_feature(app: "FastAPI", feat: LazyFeature, module: object, features: tuple[LazyFeature, ...]) -> None: before: Final = len(app.router.routes) feat.register_fn(app, module) - previous: Final[Mapping[str, tuple[BaseRoute, ...]]] = ( - app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({}) - ) lazy_routes: Final[Mapping[str, tuple[BaseRoute, ...]]] = MappingProxyType( - {**previous, feat.module_path: tuple(app.router.routes[before:])} + {**_lazy_routes(app), feat.module_path: tuple(app.router.routes[before:])} ) app.state.lazy_routes = lazy_routes # rebind-ok: the app owns the record of which routes each feature added app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table @@ -543,11 +550,8 @@ def _register_all_on_startup(inner: "Lifespan[FastAPI]", features: tuple[LazyFea def _restore_registry_order(app: "FastAPI", features: tuple[LazyFeature, ...]) -> None: present: Final = frozenset(id(route) for route in app.router.routes) - registered: Final[Mapping[str, tuple[BaseRoute, ...]]] = ( - app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({}) - ) still_routed: Final = MappingProxyType( - {module_path: tuple(r for r in routes if id(r) in present) for module_path, routes in registered.items()} + {module_path: tuple(r for r in routes if id(r) in present) for module_path, routes in _lazy_routes(app).items()} ) app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table _in_registry_order(app.router.routes, still_routed, features, _lazy_slots(app)) @@ -599,6 +603,13 @@ def _make_warmup_router(app: "FastAPI", features: tuple[LazyFeature, ...] = LAZY return router +def lazy_owned_routes(app: "FastAPI") -> frozenset[int]: + """ids of the routes lazy features have registered on this app. A route added later at + one of their paths (a config pass-through at /v1/decisions) goes ahead of them, the + precedence lazy mode gives it when the feature has not loaded by the time the config is read.""" + return frozenset(id(route) for route in chain.from_iterable(_lazy_routes(app).values())) + + def loaded_lazy_modules(app: "FastAPI") -> frozenset[str]: """The set of lazy feature modules whose routers are actually registered on this app (tracked by _install), empty until a feature loads or eager startup runs. diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 9eaf6c7e9ed..1373db220d6 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -9492,6 +9492,61 @@ } } }, + "decisions": { + "components": { + "schemas": {} + }, + "paths": { + "/decisions": { + "post": { + "operationId": "decisions_decisions_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Decisions", + "tags": [ + "decisions" + ] + } + }, + "/v1/decisions": { + "post": { + "operationId": "decisions_v1_decisions_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Decisions", + "tags": [ + "decisions" + ] + } + } + } + }, "evals": { "components": { "schemas": { diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0abec51cc49..b993a5d8d6d 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -478,6 +478,8 @@ class LiteLLMRoutes(enum.Enum): "/v1/search", "/search/{search_tool_name}", "/v1/search/{search_tool_name}", + "/decisions", + "/v1/decisions", # OCR "/ocr", "/v1/ocr", diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index 5b3930b3299..5d8287227a5 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -29,6 +29,7 @@ _MANAGED_MODEL_ROUTES: Final = frozenset( "audio/speech", "moderations", "rerank", + "decisions", "ocr", ), ) @@ -72,6 +73,7 @@ _MODEL_ROUTE_KINDS: Final[ "/audio/transcriptions": "moderation", "/audio/speech": "speech", "/rerank": "body", + "/decisions": "body", "/messages/count_tokens": "body", ":countTokens": "path", } diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 78c3f53c44f..c85169f0ba5 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -176,6 +176,7 @@ ProxyRouteType: TypeAlias = Literal[ "avector_store_file_delete", "aocr", "asearch", + "adecisions", "avideo_generation", "avideo_list", "avideo_status", @@ -1961,6 +1962,7 @@ class ProxyBaseLLMRequestProcessing: "avector_store_file_delete", "aocr", "asearch", + "adecisions", "avideo_generation", "avideo_list", "avideo_status", diff --git a/litellm/proxy/decisions_endpoints/__init__.py b/litellm/proxy/decisions_endpoints/__init__.py new file mode 100644 index 00000000000..ea9b7835485 --- /dev/null +++ b/litellm/proxy/decisions_endpoints/__init__.py @@ -0,0 +1 @@ +__all__ = () diff --git a/litellm/proxy/decisions_endpoints/endpoints.py b/litellm/proxy/decisions_endpoints/endpoints.py new file mode 100644 index 00000000000..7dba64ee791 --- /dev/null +++ b/litellm/proxy/decisions_endpoints/endpoints.py @@ -0,0 +1,102 @@ +from typing import Annotated, Final + +from fastapi import APIRouter, Depends, Request, Response +from fastapi.responses import ORJSONResponse # pyright: ignore[reportDeprecated] # required endpoint contract +from pydantic import TypeAdapter, ValidationError + +from litellm.exceptions import BadRequestError +from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.types.decisions import DecisionsRequestBody + +router: Final = APIRouter() +_REQUEST_DATA_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object]) +_DECISIONS_REQUEST_BODY_ADAPTER: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody) +_GENERAL_SETTINGS_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object]) +_OPTIONAL_STRING_ADAPTER: Final[TypeAdapter[str | None]] = TypeAdapter(str | None) +_OPTIONAL_FLOAT_ADAPTER: Final[TypeAdapter[float | None]] = TypeAdapter(float | None) + + +@router.post( + "/v1/decisions", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract + tags=["decisions"], +) +@router.post( + "/decisions", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract + tags=["decisions"], +) +async def decisions( + request: Request, + fastapi_response: Response, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +): + from litellm.proxy.proxy_server import ( + general_settings as proxy_general_settings, + ) + from litellm.proxy.proxy_server import ( + llm_router, + proxy_config, + proxy_logging_obj, + user_max_tokens, + user_request_timeout, + version, + ) + from litellm.proxy.proxy_server import ( + user_api_base as proxy_user_api_base, + ) + from litellm.proxy.proxy_server import ( + user_model as proxy_user_model, + ) + from litellm.proxy.proxy_server import ( + user_temperature as proxy_user_temperature, + ) + + data: Final = _REQUEST_DATA_ADAPTER.validate_json(await request.body()) + general_settings: Final = _GENERAL_SETTINGS_ADAPTER.validate_python(proxy_general_settings) + user_api_base: Final = _OPTIONAL_STRING_ADAPTER.validate_python(proxy_user_api_base) + user_model: Final = _OPTIONAL_STRING_ADAPTER.validate_python(proxy_user_model) + user_temperature: Final = _OPTIONAL_FLOAT_ADAPTER.validate_python(proxy_user_temperature) + processor: Final = ProxyBaseLLMRequestProcessing(data=data) + try: + _DECISIONS_REQUEST_BODY_ADAPTER.validate_python(data) + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="adecisions", + proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, + general_settings=general_settings, + proxy_config=proxy_config, + select_data_generator=None, + model=None, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + version=version, + ) + except ValidationError as error: + bad_request_error: Final = BadRequestError( + message=f"Invalid Decisions request: {error}", + model=str(data.get("model", "")), + llm_provider="", + ) + raise await processor._handle_llm_api_exception( + e=bad_request_error, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + version=version, + ) + except Exception as error: + raise await processor._handle_llm_api_exception( + e=error, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + version=version, + ) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index a5b414e5a18..f5ba7f9e877 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -28,6 +28,7 @@ from fastapi import ( from fastapi.responses import StreamingResponse from pydantic import TypeAdapter from starlette.datastructures import UploadFile as StarletteUploadFile +from starlette.routing import BaseRoute, Route from starlette.websockets import WebSocketState from websockets.asyncio.client import connect from websockets.exceptions import ( @@ -67,6 +68,7 @@ from litellm.llms.base_llm.managed_resources.utils import ( from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.llms.oss_decision import validate_oss_request from litellm.passthrough import BasePassthroughUtils +from litellm.proxy._lazy_features import lazy_owned_routes from litellm.proxy._types import ( ConfigFieldInfo, ConfigFieldUpdate, @@ -2942,35 +2944,34 @@ def _extract_model_from_vertex_ai_setup(setup_response: Mapping[str, object]) -> return None +def _placed_ahead(routes: Sequence[BaseRoute], moving: BaseRoute, before: BaseRoute) -> tuple[BaseRoute, ...]: + kept: Final = tuple(route for route in routes if route is not moving) + at: Final = next(index for index, route in enumerate(kept) if route is before) + return (*kept[:at], moving, *kept[at:]) + + class SafeRouteAdder: """ Wrapper class for adding routes to FastAPI app. - Only adds routes if they don't already exist on the app. + Only adds routes if they don't already exist on the app. A route a lazy feature registered + does not count: a route added at its path goes ahead of it, the precedence a config + pass-through at /v1/decisions gets in lazy mode, where the feature has not loaded yet. """ + @staticmethod + def _colliding_routes(app: FastAPI, path: str, methods: Sequence[str]) -> tuple[Route, ...]: + wanted: Final = frozenset(methods) + return tuple( + route + for route in app.routes + if isinstance(route, Route) and route.path == path and not wanted.isdisjoint(route.methods or ()) + ) + @staticmethod def _is_path_registered(app: FastAPI, path: str, methods: list[str]) -> bool: - """ - Check if a path with any of the specified methods is already registered on the app. - - Args: - app: The FastAPI application instance - path: The path to check (e.g., "/v1/chat/completions") - methods: List of HTTP methods to check (e.g., ["GET", "POST"]) - - Returns: - True if the path is already registered with any of the methods, False otherwise - """ - for route in app.routes: - # Use getattr to safely access route attributes - route_path = getattr(route, "path", None) - route_methods = getattr(route, "methods", None) - - if route_path == path and route_methods is not None: - # Check if any of the methods overlap - if any(method in route_methods for method in methods): - return True - return False + """True when a route the app itself defines already serves the path with one of the methods.""" + lazy_owned: Final = lazy_owned_routes(app) + return any(id(route) not in lazy_owned for route in SafeRouteAdder._colliding_routes(app, path, methods)) @staticmethod def add_api_route_if_not_exists( @@ -3001,12 +3002,17 @@ class SafeRouteAdder: ) return False + shadowed: Final = SafeRouteAdder._colliding_routes(app, path, methods) app.add_api_route( path=path, endpoint=endpoint, methods=methods, dependencies=dependencies, ) + if shadowed: + app.router.routes[:] = _placed_ahead( # rebind-ok: the app owns its route table + app.router.routes, app.router.routes[-1], shadowed[0] + ) verbose_proxy_logger.debug( "Successfully added route: %s with methods %s", path, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 7da09ddcb68..49da000155b 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -93,6 +93,7 @@ ROUTE_ENDPOINT_MAPPING: Final = { "acompact_responses": "/responses/compact", "aocr": "/ocr", "asearch": "/search", + "adecisions": "/decisions", "avideo_generation": "/videos", "avideo_list": "/videos", "avideo_status": "/videos/{video_id}", @@ -487,6 +488,7 @@ RouteType = Literal[ "avector_store_file_delete", "aocr", "asearch", + "adecisions", "avideo_generation", "avideo_list", "avideo_status", diff --git a/litellm/router.py b/litellm/router.py index 3662d1f43eb..b93a4abdf03 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -2110,6 +2110,11 @@ class Router: self.asearch = self.factory_function(asearch, call_type="asearch") self.search = self.factory_function(search, call_type="search") + from litellm.decisions import adecisions, decisions + + self.adecisions = self.factory_function(adecisions, call_type="adecisions") + self.decisions = self.factory_function(decisions, call_type="decisions") + def _initialize_video_endpoints(self): """Initialize video endpoints.""" from litellm.videos import ( @@ -6663,6 +6668,8 @@ class Router: "ocr", "asearch", "search", + "adecisions", + "decisions", "aadapter_generate_content", "avideo_generation", "video_generation", @@ -6736,6 +6743,7 @@ class Router: "generate_content_stream", "ocr", "search", + "decisions", "video_generation", "video_list", "video_status", @@ -6903,6 +6911,7 @@ class Router: "agenerate_content_stream", "aocr", "ocr", + "adecisions", "avideo_generation", "avideo_list", "avideo_status", diff --git a/litellm/types/decisions.py b/litellm/types/decisions.py new file mode 100644 index 00000000000..80db26f201d --- /dev/null +++ b/litellm/types/decisions.py @@ -0,0 +1,123 @@ +from collections.abc import Mapping, Sequence +from typing import Annotated, Literal, TypeAlias + +from pydantic import ConfigDict, Field, PrivateAttr, model_validator, with_config +from typing_extensions import ReadOnly, Required, TypedDict + +from litellm.types.llms.base import LiteLLMPydanticObjectBase + +DecisionsJSON: TypeAlias = str | Mapping[str, object] | Sequence[object] +NoulCriteria: TypeAlias = Mapping[Literal["true", "false"], DecisionsJSON | None] + + +class NoulQuestion(LiteLLMPydanticObjectBase): + type: Literal["noul"] + instructions: DecisionsJSON | None = None + criteria: NoulCriteria | None = None + + model_config = ConfigDict(extra="allow", frozen=True) + + @model_validator(mode="after") + def require_instructions_or_criteria(self) -> "NoulQuestion": + if self.instructions is None and self.criteria is None: + raise ValueError("A noul question requires instructions or criteria") + return self + + +class ChoiceQuestion(LiteLLMPydanticObjectBase): + type: Literal["choice"] + instructions: DecisionsJSON | None = None + criteria: Annotated[Mapping[str, DecisionsJSON | None], Field(min_length=1, max_length=255)] + + model_config = ConfigDict(extra="allow", frozen=True) + + +class ScoreQuestion(LiteLLMPydanticObjectBase): + type: Literal["score"] + instructions: DecisionsJSON | None = None + criteria: Annotated[Sequence[DecisionsJSON], Field(min_length=1, max_length=10)] + + model_config = ConfigDict(extra="allow", frozen=True) + + +DecisionQuestion: TypeAlias = Annotated[ + NoulQuestion | ChoiceQuestion | ScoreQuestion, + Field(discriminator="type"), +] + +DecisionQuestionMap: TypeAlias = Annotated[ + Mapping[Annotated[str, Field(min_length=1)], DecisionQuestion], + Field(min_length=1, max_length=128), +] + + +class DecisionsRequestBody(LiteLLMPydanticObjectBase): + state: DecisionsJSON + questions: DecisionQuestionMap + + model_config = ConfigDict(extra="allow", frozen=True) + + +class DecisionsRequest(DecisionsRequestBody): + model: str + + +@with_config(ConfigDict(extra="allow")) +class DecisionsCallParams(TypedDict, total=False): + model: Required[ReadOnly[str]] + state: Required[ReadOnly[DecisionsJSON]] + questions: Required[ReadOnly[DecisionQuestionMap]] + api_key: ReadOnly[str | None] + api_base: ReadOnly[str | None] + timeout: ReadOnly[float | None] + custom_llm_provider: ReadOnly[str | None] + extra_headers: ReadOnly[Mapping[str, str] | None] + + +class NoulAnswer(LiteLLMPydanticObjectBase): + type: Literal["noul"] + noul: float + + model_config = ConfigDict(extra="allow", frozen=True) + + +class ChoiceAnswer(LiteLLMPydanticObjectBase): + type: Literal["choice"] + choice: str + confidence: float + probabilities: Mapping[str, float] + + model_config = ConfigDict(extra="allow", frozen=True) + + +class ScoreAnswer(LiteLLMPydanticObjectBase): + type: Literal["score"] + score: float + confidence: float + legend: Mapping[str, DecisionsJSON] + probabilities: Mapping[str, float] + + model_config = ConfigDict(extra="allow", frozen=True) + + +DecisionAnswer: TypeAlias = Annotated[ + NoulAnswer | ChoiceAnswer | ScoreAnswer, + Field(discriminator="type"), +] + + +class DecisionsUsage(LiteLLMPydanticObjectBase): + input_tokens: int = 0 + output_tokens: int = 0 + + model_config = ConfigDict(extra="allow", frozen=True) + + +class DecisionsResponse(LiteLLMPydanticObjectBase): + model: str | None = None + answers: Mapping[str, DecisionAnswer] + usage: DecisionsUsage | None = None + + model_config = ConfigDict(extra="allow", frozen=True) + + _hidden_params: dict[str, object] = PrivateAttr(default_factory=dict) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 6919fd6fd27..a33eeaccaa3 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -468,6 +468,8 @@ class CallTypes(str, Enum): arerank = "arerank" search = "search" asearch = "asearch" + decisions = "decisions" + adecisions = "adecisions" arealtime = "_arealtime" aresponses_websocket = "_aresponses_websocket" create_batch = "create_batch" @@ -654,6 +656,8 @@ CallTypesLiteral = Literal[ "arerank", "search", "asearch", + "decisions", + "adecisions", "_arealtime", "_aresponses_websocket", "create_batch", @@ -763,6 +767,8 @@ API_ROUTE_TO_CALL_TYPES: Final[Mapping[str, Sequence[CallTypes]]] = { # Search "/search": [CallTypes.asearch, CallTypes.search], "/v1/search": [CallTypes.asearch, CallTypes.search], + "/decisions": [CallTypes.adecisions, CallTypes.decisions], + "/v1/decisions": [CallTypes.adecisions, CallTypes.decisions], # Batches "/batches": [CallTypes.acreate_batch, CallTypes.create_batch], "/v1/batches": [CallTypes.acreate_batch, CallTypes.create_batch], @@ -4048,6 +4054,8 @@ class LlmProviders(str, Enum): OLLAMA_CHAT = "ollama_chat" DEEPINFRA = "deepinfra" PERPLEXITY = "perplexity" + TYPESAFE = "typesafe" + STRANDS_DECIDER = "strands_decider" MISTRAL = "mistral" MILVUS = "milvus" GROQ = "groq" diff --git a/litellm/utils.py b/litellm/utils.py index f80a2851ba6..d72588e2b00 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1218,6 +1218,9 @@ def function_setup( if isinstance(search_query, list) else search_query ) + elif call_type in (CallTypes.decisions.value, CallTypes.adecisions.value): + decisions_state: Final = args[1] if len(args) > 1 else kwargs.get("state", "") + messages = decisions_state if isinstance(decisions_state, str) else json.dumps(decisions_state) elif call_type in (CallTypes.image_edit.value, CallTypes.aimage_edit.value): messages = args[1] if len(args) > 1 else kwargs.get("prompt") elif call_type in (CallTypes.ocr.value, CallTypes.aocr.value): diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index ea383ef4c11..3d2acf4e9d9 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -15612,6 +15612,38 @@ "prompt_cache_min_tokens": 1024, "source": "https://platform.claude.com/docs/en/about-claude/pricing" }, + "cloudflare/clef": { + "input_cost_per_token": 2.4e-07, + "litellm_provider": "cloudflare", + "max_input_tokens": 65536, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://developers.cloudflare.com/workers-ai/platform/pricing/" + }, + "cloudflare/clef-flash": { + "input_cost_per_token": 9e-08, + "litellm_provider": "cloudflare", + "max_input_tokens": 65536, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://developers.cloudflare.com/workers-ai/platform/pricing/" + }, + "cloudflare/@cf/cloudflare/clef": { + "input_cost_per_token": 2.4e-07, + "litellm_provider": "cloudflare", + "max_input_tokens": 65536, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://developers.cloudflare.com/workers-ai/platform/pricing/" + }, + "cloudflare/@cf/cloudflare/clef-flash": { + "input_cost_per_token": 9e-08, + "litellm_provider": "cloudflare", + "max_input_tokens": 65536, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://developers.cloudflare.com/workers-ai/platform/pricing/" + }, "cloudflare/@cf/meta/llama-2-7b-chat-fp16": { "input_cost_per_token": 1.923e-06, "litellm_provider": "cloudflare", @@ -44421,6 +44453,14 @@ "mode": "chat", "output_cost_per_token": 2.8e-07 }, + "perplexity/pplx-decider-v1-27b": { + "input_cost_per_token": 4e-08, + "litellm_provider": "perplexity", + "max_input_tokens": 262144, + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://docs.perplexity.ai/api-reference/decisions-post" + }, "perplexity/sonar": { "input_cost_per_token": 1e-06, "litellm_provider": "perplexity", @@ -72800,6 +72840,16 @@ "notes": "Self-hosted decision model; infrastructure costs are paid separately" } }, + "strands_decider/strands-decider-2B-hobson-v19": { + "input_cost_per_token": 0.0, + "litellm_provider": "strands_decider", + "mode": "evaluation", + "output_cost_per_token": 0.0, + "source": "https://huggingface.co/StrandsAgents/strands-decider-2B-hobson-v19", + "metadata": { + "notes": "Self-hosted decision model; infrastructure costs are paid separately" + } + }, "typesafe/jev-1.13.0": { "input_cost_per_token": 4.2e-08, "litellm_provider": "typesafe", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index eb27d3fe810..7ffaacdb3aa 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2566,6 +2566,20 @@ "image_variations": true } }, + "typesafe": { + "display_name": "TypeSafe (`typesafe`)", + "url": "https://docs.typesafe.ai/models", + "endpoints": { + "systemone": true + } + }, + "strands_decider": { + "display_name": "Strands Decider (`strands_decider`)", + "url": "https://docs.litellm.ai/docs/providers", + "endpoints": { + "systemone": true + } + }, "tavily": { "display_name": "Tavily (`tavily`)", "url": "https://docs.litellm.ai/docs/search/tavily", diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 978ac2ec092..0d199770181 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -50,12 +50,16 @@ def signal_group(group: int, action: int) -> None: pass +def graceful_stop_seconds() -> float: + return max(30.0, float(os.environ.get("INTEGRATION_PROXY_READY_SECONDS", "70"))) + + def stop_root_process(process: subprocess.Popen[bytes]) -> bool: if process.poll() is not None: return True process.terminate() try: - process.wait(timeout=30) + process.wait(timeout=graceful_stop_seconds()) except subprocess.TimeoutExpired: return False return True @@ -183,6 +187,7 @@ def owned_proxy_process( remove_environment: tuple[str, ...] = (), workers: int = 1, database_setup: tuple[str, ...] = DB_PUSH, + extra_arguments: tuple[str, ...] = (), ) -> Iterator[OwnedProxy]: root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3]) environment: Final = { @@ -209,6 +214,7 @@ def owned_proxy_process( "--num_workers", str(workers), *database_setup, + *extra_arguments, ) launch: Final = _launch_until_bound(command, root, environment, output, _PORT_ATTEMPTS) process: Final = launch.process diff --git a/tests/integration/cost_calculation/cost_tracking_case.py b/tests/integration/cost_calculation/cost_tracking_case.py index ac3fcd33d2e..8772141236e 100644 --- a/tests/integration/cost_calculation/cost_tracking_case.py +++ b/tests/integration/cost_calculation/cost_tracking_case.py @@ -246,6 +246,7 @@ class CostTrackingTestCase(BaseModel): "/v1/audio/speech", "/v1/images/generations", "/v1/images/edits", + "/v1/decisions", ] | Annotated[str, Field(pattern=r"^/(gemini|anthropic|bedrock)/")] ) = "/v1/chat/completions" diff --git a/tests/integration/cost_calculation/cost_tracking_cases.json b/tests/integration/cost_calculation/cost_tracking_cases.json index 09dfa66012f..4ec5a1c568a 100644 --- a/tests/integration/cost_calculation/cost_tracking_cases.json +++ b/tests/integration/cost_calculation/cost_tracking_cases.json @@ -57,6 +57,13 @@ "search_context_size_high": 0.012 } }, + "perplexity/pplx-decider-v1-27b": { + "input_cost_per_token": 4e-08, + "litellm_provider": "perplexity", + "max_input_tokens": 262144, + "mode": "evaluation", + "output_cost_per_token": 0.0 + }, "deepseek/deepseek-v4-chat": { "litellm_provider": "deepseek", "mode": "chat", @@ -30925,6 +30932,47 @@ "completion_tokens": 412 } }, + { + "name": "perplexity/pplx-decider-v1-27b-decisions", + "covers": "quota_management.spend_tracking.decisions_costs", + "model": "perplexity/pplx-decider-v1-27b", + "endpoint": "/v1/decisions", + "request": { + "model": "$MODEL", + "state": { + "source": "cost-tracking" + }, + "questions": { + "is_defect": { + "type": "noul", + "instructions": "Is this a defect?" + } + } + }, + "response": { + "content_type": "application/json", + "body": { + "model": "pplx-decider-v1-27b", + "answers": { + "is_defect": { + "type": "noul", + "noul": 0.9 + } + }, + "usage": { + "input_tokens": 367, + "output_tokens": 3 + } + } + }, + "expected": { + "spend": 1.468e-05, + "input_cost": 1.468e-05, + "output_cost": 0.0, + "prompt_tokens": 367, + "completion_tokens": 3 + } + }, { "name": "gpt-5.6-client_disconnect_mid_stream", "covers": "quota_management.spend_tracking.scripted_wire.client_disconnect", diff --git a/tests/integration/cost_calculation/test_cost_tracking.py b/tests/integration/cost_calculation/test_cost_tracking.py index 82878634677..cfe0a99ef4b 100644 --- a/tests/integration/cost_calculation/test_cost_tracking.py +++ b/tests/integration/cost_calculation/test_cost_tracking.py @@ -268,6 +268,31 @@ def test_case_bills_expected_cost(gateway: Gateway, case: CostTrackingTestCase) assert row.spend == 0, f"{case.name}: failure spend was {row.spend}" return assert response.is_success, f"{case.name}: proxy returned {response.status_code}: {response.text[:400]}" + if case.endpoint == "/v1/decisions": + observed: Final = JSON_OBJECT.validate_json( + httpx.get(f"{gateway.upstream_url}/__observations", timeout=5, trust_env=False).content + ) + decision_observations: Final = tuple( + value + for value in observed["requests"] + if isinstance(value, dict) and value.get("path") == f"/{scenario_id}/v1/decisions" + ) + assert decision_observations == ( + { + "path": f"/{scenario_id}/v1/decisions", + "authorization": "Bearer sk-scripted-provider", + "body": { + "model": "pplx-decider-v1-27b", + "state": {"source": "cost-tracking"}, + "questions": { + "is_defect": { + "type": "noul", + "instructions": "Is this a defect?", + } + }, + }, + }, + ) if case.response.content_type == "text/event-stream": _assert_stream_has_no_error(response.text) rows: Final = poll_rows(key, len(responses) + (prior_response_id is not None)) diff --git a/tests/integration/management/test_model_health_check.py b/tests/integration/management/test_model_health_check.py index 7d1b1cc2ea5..17bd1c2ebfe 100644 --- a/tests/integration/management/test_model_health_check.py +++ b/tests/integration/management/test_model_health_check.py @@ -1,8 +1,41 @@ +import os import uuid from typing import Final import httpx -from integration._support.client import Gateway, object_value +from integration._support.client import Gateway, object_value, string_value +from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario +from integration.cost_calculation.cost_tracking_case import JsonResponse +from pydantic import JsonValue + +_DECISIONS_PROBE_REPLY: Final = JsonResponse( + content_type="application/json", + body={ + "model": "pplx-decider-v1-27b", + "answers": {"reachable": {"type": "noul", "noul": 1.0}}, + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, +) +_CONFIGURED_PROBE_REPLY: Final = JsonResponse( + content_type="application/json", + body={ + "model": "jev-custom", + "answers": {"alive": {"type": "choice", "choice": "yes", "confidence": 0.9, "probabilities": {"yes": 0.9}}}, + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, +) +_STRANDS_PROBE_REPLY: Final = JsonResponse( + content_type="application/json", + body={ + "model": "strands-decider-2B-hobson-v19", + "answers": {"reachable": {"type": "noul", "noul": 1.0}}, + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, +) +_CONFIGURED_STATE: Final[dict[str, JsonValue]] = {"ticket": "health probe"} +_CONFIGURED_QUESTIONS: Final[dict[str, JsonValue]] = { + "alive": {"type": "choice", "criteria": {"yes": "the service answers", "no": "the service is down"}} +} def test_health_check_of_a_model_added_through_the_api_calls_its_upstream_and_reports_it_healthy( @@ -32,3 +65,85 @@ def test_health_check_of_a_model_added_through_the_api_calls_its_upstream_and_re assert [request["body"]["model"] for request in upstream.get("/__observations").json()["requests"]] == [ provider_model ] + + +def _health_report(gateway: Gateway, model: str) -> dict[str, JsonValue]: + health: Final = gateway.request("GET", "/health", params={"model": model}) + assert health.status_code == 200, health.text + return health.json() + + +def _probes_sent_to(gateway: Gateway, handle: ScenarioHandle) -> list[tuple[str, JsonValue]]: + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream: + requests: Final = upstream.get("/__observations").json()["requests"] + return [ + (string_value(request["path"]), request["body"]) + for request in map(object_value, requests) + if string_value(request["path"]).startswith(f"/{handle.scenario_id}/") + ] + + +def test_evaluation_mode_health_check_resolves_the_mode_from_the_cost_map_and_sends_the_default_probe( + gateway: Gateway, +) -> None: + with gateway.scenario() as scenario: + handle: Final = register_scenario(f"health-decisions-{uuid.uuid4().hex[:12]}", _DECISIONS_PROBE_REPLY) + scenario.cleanups.callback(delete_scenario, handle) + model: Final = scenario.model(model="perplexity/pplx-decider-v1-27b", api_base=handle.api_base()) + report: Final = _health_report(gateway, model) + assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report + assert _probes_sent_to(gateway, handle) == [ + ( + f"/{handle.scenario_id}/v1/decisions", + { + "model": "pplx-decider-v1-27b", + "state": os.environ.get("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"), + "questions": {"reachable": {"type": "noul", "instructions": "Is the service reachable?"}}, + }, + ) + ] + + +def test_evaluation_mode_health_check_sends_the_configured_state_and_questions(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = register_scenario(f"health-decisions-{uuid.uuid4().hex[:12]}", _CONFIGURED_PROBE_REPLY) + scenario.cleanups.callback(delete_scenario, handle) + model: Final = scenario.model( + model="typesafe/jev-custom", + api_base=handle.api_base(), + model_info={ + "mode": "evaluation", + "health_check_params": {"state": _CONFIGURED_STATE, "questions": _CONFIGURED_QUESTIONS}, + }, + ) + report: Final = _health_report(gateway, model) + assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report + assert _probes_sent_to(gateway, handle) == [ + ( + f"/{handle.scenario_id}/v1/systemone", + {"model": "jev-custom", "state": _CONFIGURED_STATE, "questions": _CONFIGURED_QUESTIONS}, + ) + ] + + +def test_evaluation_mode_health_check_of_the_self_hosted_strands_model_resolves_the_mode_from_the_cost_map( + gateway: Gateway, +) -> None: + with gateway.scenario() as scenario: + handle: Final = register_scenario(f"health-decisions-{uuid.uuid4().hex[:12]}", _STRANDS_PROBE_REPLY) + scenario.cleanups.callback(delete_scenario, handle) + model: Final = scenario.model( + model="strands_decider/strands-decider-2B-hobson-v19", api_base=handle.api_base(), api_key=None + ) + report: Final = _health_report(gateway, model) + assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report + assert _probes_sent_to(gateway, handle) == [ + ( + f"/{handle.scenario_id}/v1/systemone", + { + "model": "strands-decider-2B-hobson-v19", + "state": os.environ.get("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"), + "questions": {"reachable": {"type": "noul", "instructions": "Is the service reachable?"}}, + }, + ) + ] diff --git a/tests/integration/providers/test_decisions_chaos.py b/tests/integration/providers/test_decisions_chaos.py new file mode 100644 index 00000000000..2e88cdbbc7f --- /dev/null +++ b/tests/integration/providers/test_decisions_chaos.py @@ -0,0 +1,262 @@ +import asyncio +import json +import re +import signal +import threading +import uuid +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "pplx-decider-v1-27b" +_CONFIG_MODEL: Final = "decisions-chaos" +_API_KEY: Final = "synthetic-decisions-key" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_QUESTIONS: Final[dict[str, JsonValue]] = {"fine": {"type": "noul", "instructions": "Is the state fine?"}} +_ROUTES: Final = ("/v1/decisions", "/decisions") + + +@dataclass(frozen=True, slots=True) +class _Call: + route: str + marker: str + fail: bool + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + call_id: str + model_group: str + text: str + + +def _calls(count: int, *, fail: bool) -> tuple[_Call, ...]: + return tuple( + _Call( + route=_ROUTES[index % len(_ROUTES)], + marker=f"{'fail' if fail and index % 2 else 'ok'}-{uuid.uuid4().hex}", + fail=fail and index % 2 == 1, + ) + for index in range(count) + ) + + +def _marker_of(request: Request) -> str: + state: Final = _JSON_OBJECT.validate_json(request.body)["state"] + assert isinstance(state, str), request.body + return state + + +def _reply(request: Request) -> Reply: + marker: Final = _marker_of(request) + if marker.startswith("fail-"): + return Reply(status=500, body=json.dumps({"error": {"message": f"scripted outage {marker}"}}).encode()) + answer: Final = { + "model": f"model-{marker}", + "answers": {"fine": {"type": "noul", "noul": 0.5}}, + "usage": {"input_tokens": 12, "output_tokens": 1}, + } + return Reply(body=json.dumps(answer).encode()) + + +async def _send(client: httpx.AsyncClient, key: str, model: str | None, call: _Call) -> _Served: + body: Final[dict[str, JsonValue]] = { + **({"model": model} if model is not None else {}), + "state": call.marker, + "questions": _QUESTIONS, + "num_retries": 0, + } + response: Final = await client.post(call.route, json=body, headers={"Authorization": f"Bearer {key}"}) + return _Served( + call=call, + status=response.status_code, + call_id=response.headers.get("x-litellm-call-id", ""), + model_group=response.headers.get("x-litellm-model-group", ""), + text=response.text, + ) + + +async def _burst( + base_url: str, key: str, model: str | None, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False +) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client: + results: Final = await asyncio.gather( + *(_send(client, key, model, call) for call in calls), return_exceptions=tolerate_transport_errors + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Served)) + + +def _assert_served_its_own(served: _Served) -> None: + assert served.call_id, served.text + if served.call.fail: + assert served.status == 500, (served.status, served.text) + assert served.call.marker in served.text, served.text + return + assert served.status == 200, (served.status, served.text) + assert _JSON_OBJECT.validate_json(served.text)["model"] == f"model-{served.call.marker}", served.text + + +def _statuses_by_call_id(call_ids: tuple[str, ...]) -> dict[str, JsonValue]: + placeholders: Final = ", ".join("%s" for _ in call_ids) + rows: Final = eventually( + lambda: read_rows( + f'SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE request_id IN ({placeholders})', call_ids + ), + lambda found: len(found) >= len(call_ids), + seconds=70, + ) + assert len(rows) == len(call_ids), rows + return {str(row["request_id"]): row["status"] for row in rows} + + +def _expected_statuses(served: tuple[_Served, ...]) -> dict[str, JsonValue]: + return {item.call_id: "failure" if item.call.fail else "success" for item in served} + + +def _health(gateway: Gateway, model: str) -> tuple[int, int]: + health: Final = gateway.request("GET", "/health", params={"model": model}) + assert health.status_code in (200, 503), health.text + report: Final = health.json() + return (report["healthy_count"], report["unhealthy_count"]) + + +def _chaos_config(wire: Wire, tmp_path: Path) -> Path: + config: Final = { + **_JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())), + "model_list": [ + { + "model_name": _CONFIG_MODEL, + "litellm_params": {"model": f"perplexity/{_BACKEND}", "api_base": wire.url, "api_key": _API_KEY}, + } + ], + } + path: Final = tmp_path / "decisions-chaos.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _worker_startups(log: Path) -> tuple[tuple[int, ...], int]: + text: Final = log.read_text() + return tuple(int(pid) for pid in _STARTED_WORKER.findall(text)), text.count("Application startup complete.") + + +def _open_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +@pytest.mark.timeout(180) +async def test_burst_over_both_routes_bills_each_call_once_with_its_own_status(gateway: Gateway) -> None: + calls: Final = _calls(30, fail=True) + with wire_server(_reply) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"perplexity/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls) + assert len(served) == 30 + for item in served: + _assert_served_its_own(item) + assert len({item.call_id for item in served}) == 30 + assert _statuses_by_call_id(tuple(item.call_id for item in served)) == _expected_statuses(served) + received: Final = wire.drain() + assert sorted(_marker_of(request) for request in received) == sorted(call.marker for call in calls) + assert {request.target for request in received} == {"/v1/decisions"}, received + + +@pytest.mark.timeout(420) +async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_the_default_model( + gateway: Gateway, tmp_path: Path +) -> None: + calls: Final = _calls(20, fail=False) + release: Final = threading.Event() + held_markers: Final[SimpleQueue[str]] = SimpleQueue() + + def held(request: Request) -> Reply: + held_markers.put(_marker_of(request)) + assert release.wait(timeout=60), "The burst was never released" + return _reply(request) + + with wire_server(held) as wire: + path: Final = _chaos_config(wire, tmp_path) + with owned_proxy_process( + gateway, tmp_path, {}, config=path, workers=2, extra_arguments=("--model", _CONFIG_MODEL) + ) as owned: + candidate: Final = owned.gateway + base_url: Final = str(candidate.client.base_url) + workers, _ = eventually( + lambda: _worker_startups(owned.log), + lambda found: len(found[0]) == 2 and found[1] == 2, + seconds=120, + ) + burst: Final = asyncio.create_task( + _burst(base_url, candidate.key, None, calls, tolerate_transport_errors=True) + ) + await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == 20, 60) + held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + served: Final = await burst + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + follow_up: Final = _Call(route="/decisions", marker=f"ok-{uuid.uuid4().hex}", fail=False) + (answered,) = await _burst(base_url, candidate.key, None, (follow_up,)) + await asyncio.to_thread( + eventually, + lambda: _worker_startups(owned.log), + lambda found: len(found[0]) == 3 and found[1] == 3, + 180, + ) + for item in (*served, answered): + _assert_served_its_own(item) + assert item.model_group == _CONFIG_MODEL, item.model_group + call_ids: Final = tuple(item.call_id for item in (*served, answered)) + assert set(_statuses_by_call_id(call_ids).values()) == {"success"} + assert len({_marker_of(request) for request in wire.drain()}) == 21 + + +@pytest.mark.timeout(180) +async def test_upstream_outage_fails_its_calls_and_recovery_on_the_same_port_restores_them(gateway: Gateway) -> None: + base_url: Final = str(gateway.client.base_url) + with gateway.scenario() as scenario: + with wire_server(_reply) as wire: + port: Final = urlsplit(wire.url).port + assert port is not None + model: Final = scenario.model(model=f"perplexity/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + before: Final = await _burst(base_url, gateway.key, model, _calls(5, fail=False)) + assert _health(gateway, model) == (1, 0) + during: Final = await _burst(base_url, gateway.key, model, _calls(5, fail=False)) + assert [item.status for item in during] == [500] * 5, [item.text for item in during] + assert _health(gateway, model) == (0, 1) + with wire_server(_reply, port=port): + after: Final = await _burst(base_url, gateway.key, model, _calls(5, fail=False)) + assert _health(gateway, model) == (1, 0) + for item in (*before, *after): + _assert_served_its_own(item) + statuses: Final = _statuses_by_call_id(tuple(item.call_id for item in (*before, *during, *after))) + assert statuses == { + **{item.call_id: "success" for item in (*before, *after)}, + **{item.call_id: "failure" for item in during}, + } diff --git a/tests/integration/providers/test_decisions_wire.py b/tests/integration/providers/test_decisions_wire.py new file mode 100644 index 00000000000..9878559cb07 --- /dev/null +++ b/tests/integration/providers/test_decisions_wire.py @@ -0,0 +1,456 @@ +import json +import math +import socket +import uuid +from collections.abc import Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy_process +from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario +from integration.cost_calculation.cost_tracking_case import JsonResponse +from pydantic import JsonValue, TypeAdapter + +import litellm + +_API_KEY: Final = "synthetic-decisions-key" +_ENV_KEY: Final = "synthetic-decisions-env-key" +_PASS_THROUGH_MODEL: Final = "gpt-6-luna" +_PASS_THROUGH_AUTHORIZATION: Final = "Bearer customer-held-upstream-key" +_PASS_THROUGH_NEIGHBOUR: Final = "decisions-beside-a-pass-through" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_USAGE: Final[dict[str, JsonValue]] = {"input_tokens": 367, "output_tokens": 3} +_STATE: Final[dict[str, JsonValue]] = {"ticket": "The export job hangs at 99%", "component": "billing"} +_QUESTIONS: Final[dict[str, JsonValue]] = { + "defect": {"type": "noul", "instructions": "Is this a defect?"}, + "severity": {"type": "choice", "criteria": {"low": "cosmetic", "high": "blocks users"}, "weight": 2}, + "confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]}, +} +_ANSWERS: Final[dict[str, JsonValue]] = { + "defect": {"type": "noul", "noul": 0.93}, + "severity": {"type": "choice", "choice": "high", "confidence": 0.8, "probabilities": {"low": 0.2, "high": 0.8}}, + "confidence": { + "type": "score", + "score": 1.0, + "confidence": 0.7, + "legend": {"0": "unsure", "1": "sure"}, + "probabilities": {"0": 0.3, "1": 0.7}, + }, +} +_CHAT_BODY: Final[dict[str, JsonValue]] = {"messages": [{"role": "user", "content": "hi"}]} +_CHAT_REPLY: Final[dict[str, JsonValue]] = { + "id": "chatcmpl-decisions-parity", + "object": "chat.completion", + "created": 1700000000, + "model": "pplx-decider-v1-27b", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4}, +} +_SPEND_QUERY: Final = ( + "SELECT spend, status, call_type, model_group, custom_llm_provider, api_base, prompt_tokens, completion_tokens, " + 'request_tags FROM "LiteLLM_SpendLogs" WHERE request_id = %s' +) + + +@dataclass(frozen=True, slots=True) +class _Provider: + name: str + model: str + path: str + body_model: str + api_key: str | None + wraps_result: bool + cost_map_key: str | None + + +_PROVIDERS: Final = ( + _Provider( + "perplexity", + "perplexity/pplx-decider-v1-27b", + "/v1/decisions", + "pplx-decider-v1-27b", + _API_KEY, + False, + "perplexity/pplx-decider-v1-27b", + ), + _Provider("typesafe", "typesafe/jev-1.13.0", "/v1/systemone", "jev-1.13.0", _API_KEY, False, "typesafe/jev-1.13.0"), + _Provider( + "openrouter", + "openrouter/typesafe/jev-1.13", + "/alpha/decisions", + "typesafe/jev-1.13", + _API_KEY, + False, + "openrouter/typesafe/jev-1.13", + ), + _Provider( + "strands_decider", "strands_decider/systemone-decider", "/v1/systemone", "systemone-decider", None, False, None + ), + _Provider( + "cloudflare", + "cloudflare/clef", + "/ai/run/@cf/cloudflare/clef", + "clef", + _API_KEY, + True, + "cloudflare/@cf/cloudflare/clef", + ), +) +_PERPLEXITY: Final = _PROVIDERS[0] +_INVALID_BODIES: Final[tuple[tuple[str, dict[str, JsonValue]], ...]] = ( + ("missing questions", {"state": _STATE}), + ("missing state", {"questions": _QUESTIONS}), + ("numeric state", {"state": 5, "questions": _QUESTIONS}), + ("empty questions", {"state": _STATE, "questions": {}}), + ("noul without instructions or criteria", {"state": _STATE, "questions": {"q": {"type": "noul"}}}), + ("choice without criteria", {"state": _STATE, "questions": {"q": {"type": "choice", "criteria": {}}}}), + ( + "score with eleven criteria", + {"state": _STATE, "questions": {"q": {"type": "score", "criteria": [f"level-{index}" for index in range(11)]}}}, + ), + ("unknown question type", {"state": _STATE, "questions": {"q": {"type": "ranking", "criteria": ["a"]}}}), +) + + +def _number(value: JsonValue) -> float: + assert isinstance(value, (int, float)) and not isinstance(value, bool), value + return float(value) + + +def _expected_spend(cost_map_key: str | None) -> float: + if cost_map_key is None: + return 0.0 + prices: Final = object_value(json.loads(Path("model_prices_and_context_window.json").read_text())[cost_map_key]) + return _number(_USAGE["input_tokens"]) * _number(prices["input_cost_per_token"]) + _number( + _USAGE["output_tokens"] + ) * _number(prices["output_cost_per_token"]) + + +def _answer_body(provider: _Provider) -> dict[str, JsonValue]: + answer: Final[dict[str, JsonValue]] = {"model": provider.body_model, "answers": _ANSWERS, "usage": _USAGE} + return {"result": answer, "success": True} if provider.wraps_result else answer + + +def _register(scenario: Scenario, body: dict[str, JsonValue], *, status: int = 200) -> ScenarioHandle: + handle: Final = register_scenario( + f"decisions-{uuid.uuid4().hex[:12]}", JsonResponse(content_type="application/json", body=body, status=status) + ) + scenario.cleanups.callback(delete_scenario, handle) + return handle + + +def _deployment(scenario: Scenario, handle: ScenarioHandle, provider: _Provider) -> str: + return scenario.model(model=provider.model, api_base=handle.api_base(), api_key=provider.api_key) + + +def _decide(gateway: Gateway, model: str, *, key: str | None = None, **extra: JsonValue) -> httpx.Response: + return gateway.request( + "POST", "/v1/decisions", {"model": model, "state": _STATE, "questions": _QUESTIONS, **extra}, key=key + ) + + +def _chat(gateway: Gateway, model: str, *, key: str | None = None, **extra: JsonValue) -> httpx.Response: + return gateway.request("POST", "/v1/chat/completions", {"model": model, **_CHAT_BODY, **extra}, key=key) + + +def _observed_requests(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]: + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream: + return tuple(map(object_value, upstream.get("/__observations").json()["requests"])) + + +def _calls_to(requests: Sequence[dict[str, JsonValue]], handle: ScenarioHandle) -> list[dict[str, JsonValue]]: + return [request for request in requests if string_value(request["path"]).startswith(f"/{handle.scenario_id}/")] + + +def _upstream_calls(gateway: Gateway, handle: ScenarioHandle) -> list[dict[str, JsonValue]]: + return _calls_to(_observed_requests(gateway), handle) + + +def _spend_row(call_id: str) -> dict[str, JsonValue]: + rows: Final = eventually(lambda: read_rows(_SPEND_QUERY, (call_id,)), lambda found: len(found) == 1, seconds=70) + return rows[0] + + +def _free_closed_port() -> int: + with socket.socket() as probe: + probe.bind(("127.0.0.1", 0)) + return probe.getsockname()[1] + + +def _pass_through_config(directory: Path, pass_through_target: str, native_api_base: str) -> Path: + base: Final = _JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + config: Final = { + **base, + "general_settings": { + **object_value(base["general_settings"]), + "pass_through_endpoints": [ + { + "path": "/v1/decisions", + "target": pass_through_target, + "headers": {"Authorization": _PASS_THROUGH_AUTHORIZATION}, + } + ], + }, + "model_list": [ + { + "model_name": _PASS_THROUGH_NEIGHBOUR, + "litellm_params": {"model": _PERPLEXITY.model, "api_base": native_api_base, "api_key": _API_KEY}, + } + ], + } + path: Final = directory / "decisions-pass-through.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@pytest.mark.parametrize("provider", _PROVIDERS, ids=lambda provider: provider.name) +def test_each_provider_gets_its_own_path_key_and_body_and_is_billed_from_the_cost_map( + gateway: Gateway, provider: _Provider +) -> None: + expected_spend: Final = _expected_spend(provider.cost_map_key) + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _answer_body(provider)) + model: Final = _deployment(scenario, handle, provider) + response: Final = _decide(gateway, model) + assert response.status_code == 200, response.text + assert response.json() == {"model": provider.body_model, "answers": _ANSWERS, "usage": _USAGE} + assert response.headers["x-litellm-model-group"] == model + assert math.isclose(float(response.headers.get("x-litellm-response-cost", "0")), expected_spend, rel_tol=1e-9) + (call,) = _upstream_calls(gateway, handle) + assert call["path"] == f"/{handle.scenario_id}{provider.path}" + assert call["authorization"] == (f"Bearer {provider.api_key}" if provider.api_key else "") + assert call["body"] == {"model": provider.body_model, "state": _STATE, "questions": _QUESTIONS} + row: Final = _spend_row(response.headers["x-litellm-call-id"]) + assert ( + row["status"], + row["call_type"], + row["custom_llm_provider"], + row["model_group"], + row["api_base"], + row["prompt_tokens"], + row["completion_tokens"], + ) == ("success", "adecisions", provider.name, model, f"{handle.api_base()}{provider.path}", 367, 3) + assert math.isclose(_number(row["spend"]), expected_spend, rel_tol=1e-9), row + + +def test_repeated_identical_requests_each_reach_the_upstream_and_are_each_billed(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) + model: Final = _deployment(scenario, handle, _PERPLEXITY) + responses: Final = tuple(_decide(gateway, model) for _ in range(2)) + assert [response.status_code for response in responses] == [200, 200], [r.text for r in responses] + call_ids: Final = tuple(response.headers["x-litellm-call-id"] for response in responses) + assert len(set(call_ids)) == 2, call_ids + assert len(_upstream_calls(gateway, handle)) == 2 + for call_id in call_ids: + assert _spend_row(call_id)["status"] == "success" + + +async def test_sdk_sync_and_async_clients_send_the_same_request(gateway: Gateway) -> None: + provider: Final = _PROVIDERS[1] + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _answer_body(provider)) + synchronous: Final = litellm.decisions( + model=provider.model, state=_STATE, questions=_QUESTIONS, api_base=handle.api_base(), api_key=_API_KEY + ) + asynchronous: Final = await litellm.adecisions( + model=provider.model, state=_STATE, questions=_QUESTIONS, api_base=handle.api_base(), api_key=_API_KEY + ) + for response in (synchronous, asynchronous): + assert response.model_dump(mode="json") == { + "model": provider.body_model, + "answers": _ANSWERS, + "usage": _USAGE, + } + calls: Final = _upstream_calls(gateway, handle) + assert len(calls) == 2, calls + for call in calls: + assert call["path"] == f"/{handle.scenario_id}{provider.path}" + assert call["authorization"] == f"Bearer {_API_KEY}" + assert call["body"] == {"model": provider.body_model, "state": _STATE, "questions": _QUESTIONS} + + +def test_gateway_only_fields_stay_at_the_gateway_and_tags_reach_the_spend_log(gateway: Gateway) -> None: + tag: Final = f"decisions-audit-{uuid.uuid4().hex[:8]}" + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) + model: Final = _deployment(scenario, handle, _PERPLEXITY) + response: Final = _decide( + gateway, model, user="auditor", num_retries=0, temperature=0.2, metadata={"tags": [tag]} + ) + assert response.status_code == 200, response.text + (call,) = _upstream_calls(gateway, handle) + assert call["body"] == {"model": _PERPLEXITY.body_model, "state": _STATE, "questions": _QUESTIONS} + row: Final = _spend_row(response.headers["x-litellm-call-id"]) + tags: Final = row["request_tags"] + assert isinstance(tags, list) and tag in tags, row + + +def test_invalid_bodies_are_refused_at_the_gateway_without_an_upstream_call(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) + model: Final = _deployment(scenario, handle, _PERPLEXITY) + for label, body in _INVALID_BODIES: + response: Final = gateway.request("POST", "/v1/decisions", {"model": model, **body}) + assert response.status_code == 400, (label, response.text) + assert "Invalid Decisions request" in response.text, (label, response.text) + assert _upstream_calls(gateway, handle) == [] + + +def test_unknown_model_is_refused_like_chat(gateway: Gateway) -> None: + model: Final = f"missing-{uuid.uuid4().hex}" + decisions: Final = _decide(gateway, model) + chat: Final = _chat(gateway, model) + assert 400 <= decisions.status_code < 500, decisions.text + assert decisions.status_code == chat.status_code, (decisions.text, chat.text) + + +def test_key_checks_match_chat(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) + model: Final = _deployment(scenario, handle, _PERPLEXITY) + anonymous: Final = gateway.client.post( + "/v1/decisions", json={"model": model, "state": _STATE, "questions": _QUESTIONS} + ) + assert anonymous.status_code == 401, anonymous.text + restricted: Final = scenario.key(models=[f"other-{uuid.uuid4().hex}"]) + refused: Final = _decide(gateway, model, key=restricted) + assert 400 <= refused.status_code < 500, refused.text + assert refused.status_code == _chat(gateway, model, key=restricted).status_code, refused.text + assert _upstream_calls(gateway, handle) == [] + spender: Final = scenario.key(max_budget=1e-06) + first: Final = _decide(gateway, model, key=spender) + assert first.status_code == 200, first.text + blocked: Final = eventually( + lambda: _decide(gateway, model, key=spender), lambda response: response.status_code != 200, seconds=70 + ) + assert 400 <= blocked.status_code < 500, blocked.text + assert blocked.status_code == _chat(gateway, model, key=spender).status_code, blocked.text + + +def test_request_body_api_base_is_refused_like_chat_without_an_upstream_call(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) + model: Final = _deployment(scenario, handle, _PERPLEXITY) + decisions: Final = _decide(gateway, model, api_base=f"http://127.0.0.1:{_free_closed_port()}") + chat: Final = _chat(gateway, model, api_base=f"http://127.0.0.1:{_free_closed_port()}") + assert 400 <= decisions.status_code < 500, decisions.text + assert decisions.status_code == chat.status_code, (decisions.text, chat.text) + assert _upstream_calls(gateway, handle) == [] + + +def test_a_deployment_without_a_key_sends_the_provider_env_key_to_its_configured_api_base(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) + model: Final = scenario.model(model=_PERPLEXITY.model, api_base=handle.api_base(), api_key=None) + response: Final = _decide(gateway, model) + assert response.status_code == 200, response.text + (call,) = _upstream_calls(gateway, handle) + assert (call["path"], call["authorization"]) == (f"/{handle.scenario_id}/v1/decisions", f"Bearer {_ENV_KEY}") + + +def test_a_deployment_opted_into_client_api_base_sends_decisions_and_chat_to_the_body_api_base( + gateway: Gateway, +) -> None: + with gateway.scenario() as scenario: + configured: Final = _register(scenario, _answer_body(_PERPLEXITY)) + decisions_target: Final = _register(scenario, _answer_body(_PERPLEXITY)) + chat_target: Final = _register(scenario, _CHAT_REPLY) + model: Final = scenario.model( + model=_PERPLEXITY.model, + api_base=configured.api_base(), + api_key=_API_KEY, + configurable_clientside_auth_params=["api_base"], + ) + decisions: Final = _decide(gateway, model, api_base=decisions_target.api_base()) + chat: Final = _chat(gateway, model, api_base=chat_target.api_base()) + assert decisions.status_code == 200, decisions.text + assert chat.status_code == 200, chat.text + observed: Final = _observed_requests(gateway) + assert [call["path"] for call in _calls_to(observed, decisions_target)] == [ + f"/{decisions_target.scenario_id}/v1/decisions" + ] + assert [call["path"] for call in _calls_to(observed, chat_target)] == [ + f"/{chat_target.scenario_id}/chat/completions" + ] + assert _calls_to(observed, configured) == [] + + +def test_a_config_pass_through_at_v1_decisions_keeps_answering_and_the_native_api_serves_decisions( + gateway: Gateway, tmp_path: Path +) -> None: + with gateway.scenario() as scenario: + pass_through_target: Final = _register( + scenario, {"model": _PASS_THROUGH_MODEL, "answers": _ANSWERS, "usage": _USAGE} + ) + native_target: Final = _register(scenario, _answer_body(_PERPLEXITY)) + config: Final = _pass_through_config( + tmp_path, f"{pass_through_target.api_base()}/v1/decisions", native_target.api_base() + ) + with owned_proxy_process(gateway, tmp_path, {}, config=config) as owned: + through: Final = _decide(owned.gateway, _PASS_THROUGH_MODEL) + native: Final = owned.gateway.request( + "POST", "/decisions", {"model": _PASS_THROUGH_NEIGHBOUR, "state": _STATE, "questions": _QUESTIONS} + ) + assert through.status_code == 200, through.text + assert through.json() == {"model": _PASS_THROUGH_MODEL, "answers": _ANSWERS, "usage": _USAGE} + observed: Final = _observed_requests(gateway) + (forwarded,) = _calls_to(observed, pass_through_target) + assert (forwarded["path"], forwarded["authorization"], object_value(forwarded["body"])["model"]) == ( + f"/{pass_through_target.scenario_id}/v1/decisions", + _PASS_THROUGH_AUTHORIZATION, + _PASS_THROUGH_MODEL, + ) + assert native.status_code == 200, native.text + assert [call["path"] for call in _calls_to(observed, native_target)] == [ + f"/{native_target.scenario_id}/v1/decisions" + ] + + +@pytest.mark.parametrize("status", (401, 429, 500)) +def test_upstream_errors_keep_their_status_and_log_an_unbilled_failure(gateway: Gateway, status: int) -> None: + marker: Final = f"scripted-{status}-{uuid.uuid4().hex[:8]}" + with gateway.scenario() as scenario: + handle: Final = _register(scenario, {"error": {"message": marker}}, status=status) + model: Final = _deployment(scenario, handle, _PERPLEXITY) + response: Final = _decide(gateway, model, num_retries=0) + assert response.status_code == status, response.text + assert marker in response.text + row: Final = _spend_row(response.headers["x-litellm-call-id"]) + assert (row["status"], row["call_type"], row["model_group"], _number(row["spend"])) == ( + "failure", + "adecisions", + model, + 0.0, + ) + + +def test_upstream_success_without_answers_is_a_gateway_side_server_error(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, {"model": _PERPLEXITY.body_model, "usage": _USAGE}) + model: Final = _deployment(scenario, handle, _PERPLEXITY) + response: Final = _decide(gateway, model, num_retries=0) + assert 500 <= response.status_code < 600, response.text + assert "answers" in response.text + assert _spend_row(response.headers["x-litellm-call-id"])["status"] == "failure" + + +def test_unreachable_upstream_fails_only_its_own_deployment(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _register(scenario, _answer_body(_PERPLEXITY)) + healthy: Final = _deployment(scenario, handle, _PERPLEXITY) + dead: Final = scenario.model( + model=_PERPLEXITY.model, api_base=f"http://127.0.0.1:{_free_closed_port()}", api_key=_API_KEY + ) + failed: Final = _decide(gateway, dead, num_retries=0) + assert 500 <= failed.status_code < 600, failed.text + assert _spend_row(failed.headers["x-litellm-call-id"])["status"] == "failure" + served: Final = _decide(gateway, healthy) + assert served.status_code == 200, served.text + assert len(_upstream_calls(gateway, handle)) == 1 diff --git a/tests/integration/proxy_config.yaml b/tests/integration/proxy_config.yaml index a3d29e42120..085adea81ac 100644 --- a/tests/integration/proxy_config.yaml +++ b/tests/integration/proxy_config.yaml @@ -23,3 +23,5 @@ vector_store_registry: api_base: os.environ/INTEGRATION_UPSTREAM_URL api_key: integration-provider-key vector_store_description: declared in tests/integration/proxy_config.yaml +environment_variables: + PERPLEXITYAI_API_KEY: synthetic-decisions-env-key diff --git a/tests/unit/decisions/__init__.py b/tests/unit/decisions/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/decisions/test_main.py b/tests/unit/decisions/test_main.py new file mode 100644 index 00000000000..106710d328f --- /dev/null +++ b/tests/unit/decisions/test_main.py @@ -0,0 +1,592 @@ +from __future__ import annotations + +import asyncio +import json +from collections.abc import Mapping +from types import MappingProxyType +from typing import Final + +import pytest +import respx + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER +from litellm.types.decisions import ( + ChoiceAnswer, + DecisionsResponse, + DecisionsUsage, + NoulAnswer, + ScoreAnswer, +) + +_QUESTIONS: Final[Mapping[str, object]] = MappingProxyType( + { + "is_defect": {"type": "noul", "instructions": "Is this a defect?", "provider_field": "kept"}, + "sentiment": {"type": "choice", "criteria": {"positive": None, "negative": "unhappy"}}, + "severity": {"type": "score", "criteria": ["none", "low", "high"]}, + } +) +_INPUT_TOKENS: Final[int] = 367 +_OUTPUT_TOKENS: Final[int] = 3 +_RESPONSE: Final[Mapping[str, object]] = { + "model": "jev-1.13", + "answers": { + "is_defect": {"type": "noul", "noul": 0.9}, + "sentiment": { + "type": "choice", + "choice": "positive", + "confidence": 0.8, + "probabilities": {"positive": 0.8, "negative": 0.2}, + }, + "severity": { + "type": "score", + "score": 1, + "confidence": 0.7, + "legend": {"0": "none", "1": "low", "2": "high"}, + "probabilities": {"0": 0.1, "1": 0.8, "2": 0.1}, + }, + }, + "usage": {"input_tokens": _INPUT_TOKENS, "output_tokens": _OUTPUT_TOKENS}, +} +_STRANDS_RESPONSE: Final[Mapping[str, object]] = { + "model": "strands-decider-2B-hobson-v19", + "answers": { + "severity": { + "type": "score", + "score": 1, + "confidence": 0.7, + "legend": {"0": "none", "1": "low", "2": "high"}, + "probabilities": {"0": 0.1, "1": 0.8, "2": 0.1}, + } + }, + "usage": {"input_tokens": 216, "output_tokens": 3}, + "latency_ms": 3722.17, +} +_PROVIDERS: Final[tuple[tuple[str, str, str, str], ...]] = ( + ( + "perplexity", + "perplexity/pplx-decider-v1-27b", + "https://api.perplexity.ai/v1/decisions", + "pplx-decider-v1-27b", + ), + ("typesafe", "typesafe/jev-1.13", "https://api.typesafe.ai/v1/systemone", "jev-1.13"), + ( + "openrouter", + "openrouter/typesafe/jev-1.13", + "https://openrouter.ai/api/alpha/decisions", + "typesafe/jev-1.13", + ), +) + + +class _RecordingLogger(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.standard_logging_object: Mapping[str, object] | None = None + + async def async_log_success_event( + self, + kwargs: Mapping[str, object], + response_obj: object, + start_time: object, + end_time: object, + ) -> None: + standard_logging_object: Final = kwargs.get("standard_logging_object") + if isinstance(standard_logging_object, dict): + self.standard_logging_object = standard_logging_object + + +async def _drain_logging_worker() -> None: + await asyncio.sleep(0) + GLOBAL_LOGGING_WORKER.start() + await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0) + + +@pytest.fixture(autouse=True) +def _httpx_transport(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("provider", "model", "url", "upstream_model"), _PROVIDERS) +async def test_adecisions_sends_the_provider_wire_contract( + provider: str, + model: str, + url: str, + upstream_model: str, + respx_mock: respx.MockRouter, +) -> None: + route: Final = respx_mock.post(url).respond(json=_RESPONSE) + + response: Final = await litellm.adecisions( + model=model, + state={"source": "unit-test"}, + questions=_QUESTIONS, + api_key="caller-key", + extra_headers={ + "x-request-tag": "decisions-test", + "AUTHORIZATION": "attacker-key", + "Content-Type": "text/plain", + }, + internal_kwarg="must-not-leak", + ) + + assert route.called + assert len(respx_mock.calls) == 1 + request: Final = respx_mock.calls[0].request + assert request.headers["authorization"] == "Bearer caller-key" + assert request.headers["content-type"] == "application/json" + assert request.headers["x-request-tag"] == "decisions-test" + assert json.loads(request.content) == { + "model": upstream_model, + "state": {"source": "unit-test"}, + "questions": { + "is_defect": { + "type": "noul", + "instructions": "Is this a defect?", + "provider_field": "kept", + }, + "sentiment": {"type": "choice", "criteria": {"positive": None, "negative": "unhappy"}}, + "severity": {"type": "score", "criteria": ["none", "low", "high"]}, + }, + } + assert isinstance(response.answers["is_defect"], NoulAnswer) + assert isinstance(response.answers["sentiment"], ChoiceAnswer) + assert isinstance(response.answers["severity"], ScoreAnswer) + assert response._hidden_params["custom_llm_provider"] == provider + + +@pytest.mark.asyncio +async def test_router_dispatches_typesafe_decisions_without_api_base( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("TYPESAFE_API_KEY", raising=False) + monkeypatch.delenv("TYPESAFE_API_BASE", raising=False) + provider_resolution: Final = litellm.get_llm_provider("typesafe/jev-latest") + + assert provider_resolution[:2] == ("jev-latest", "typesafe") + + router: Final = litellm.Router( + model_list=[ + { + "model_name": "jev", + "litellm_params": { + "model": "typesafe/jev-latest", + "api_key": "k", + }, + } + ] + ) + upstream: Final = respx_mock.post("https://api.typesafe.ai/v1/systemone").respond(json=_RESPONSE) + + response: Final = await router.adecisions( + model="jev", + state="router-test", + questions={ + "sentiment": { + "type": "choice", + "criteria": {"positive": None, "negative": "unhappy"}, + } + }, + ) + + assert upstream.called + assert len(respx_mock.calls) == 1 + assert json.loads(respx_mock.calls[0].request.content) == { + "model": "jev-latest", + "state": "router-test", + "questions": { + "sentiment": { + "type": "choice", + "criteria": {"positive": None, "negative": "unhappy"}, + } + }, + } + assert respx_mock.calls[0].request.headers["authorization"] == "Bearer k" + assert isinstance(response.answers["sentiment"], ChoiceAnswer) + assert response.answers["sentiment"].choice == "positive" + + +def test_decisions_uses_the_same_wire_contract_for_sync_calls(respx_mock: respx.MockRouter) -> None: + route: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + + response: Final = litellm.decisions( + model="perplexity/pplx-decider-v1-27b", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + api_key="caller-key", + ) + + assert route.called + assert response.model == "jev-1.13" + + +def test_openrouter_response_keeps_provider_fields(respx_mock: respx.MockRouter) -> None: + payload: Final = { + **_RESPONSE, + "id": "decision-1", + "provider": "typesafe", + "usage": {**_RESPONSE["usage"], "cost": 0.25}, + } + respx_mock.post("https://openrouter.ai/api/alpha/decisions").respond(json=payload) + + response: Final = litellm.decisions( + model="openrouter/typesafe/jev-1.13", + state="review", + questions=_QUESTIONS, + api_key="caller-key", + ) + + assert response.model_extra["id"] == "decision-1" + assert response.model_extra["provider"] == "typesafe" + assert response.usage is not None + assert response.usage.model_extra["cost"] == 0.25 + + +def test_decisions_cost_uses_litellm_token_pricing() -> None: + response: Final = DecisionsResponse( + model="pplx-decider-v1-27b", + answers={}, + usage=DecisionsUsage(input_tokens=_INPUT_TOKENS, output_tokens=_OUTPUT_TOKENS), + ) + response._hidden_params = { + "model": "perplexity/pplx-decider-v1-27b", + "custom_llm_provider": "perplexity", + } + + cost: Final = litellm.completion_cost(completion_response=response) + perplexity_cost: Final = litellm.model_cost["perplexity/pplx-decider-v1-27b"] + expected_cost: Final = _INPUT_TOKENS * float(perplexity_cost["input_cost_per_token"]) + _OUTPUT_TOKENS * float( + perplexity_cost["output_cost_per_token"] + ) + + assert expected_cost > 0 + assert cost == pytest.approx(expected_cost) + + +@pytest.mark.asyncio +async def test_decisions_cost_is_in_standard_logging_object(respx_mock: respx.MockRouter) -> None: + respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + recording_logger: Final = _RecordingLogger() + original_callbacks: Final = litellm.callbacks + litellm.callbacks = [recording_logger] + + try: + await litellm.adecisions( + model="perplexity/pplx-decider-v1-27b", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + api_key="caller-key", + ) + await _drain_logging_worker() + finally: + litellm.callbacks = original_callbacks + + assert recording_logger.standard_logging_object is not None + perplexity_cost: Final = litellm.model_cost["perplexity/pplx-decider-v1-27b"] + expected_cost: Final = _INPUT_TOKENS * float(perplexity_cost["input_cost_per_token"]) + _OUTPUT_TOKENS * float( + perplexity_cost["output_cost_per_token"] + ) + + assert expected_cost > 0 + assert recording_logger.standard_logging_object["response_cost"] == pytest.approx(expected_cost) + assert recording_logger.standard_logging_object["prompt_tokens"] == _INPUT_TOKENS + assert recording_logger.standard_logging_object["completion_tokens"] == _OUTPUT_TOKENS + + +@pytest.mark.asyncio +async def test_unknown_provider_is_rejected_before_http(respx_mock: respx.MockRouter) -> None: + with pytest.raises(litellm.BadRequestError, match="Supported providers"): + await litellm.adecisions( + model="unknown/jev-1.13", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + api_key="caller-key", + ) + + assert len(respx_mock.calls) == 0 + + +@pytest.mark.asyncio +async def test_empty_custom_provider_is_rejected_before_http(respx_mock: respx.MockRouter) -> None: + with pytest.raises(litellm.BadRequestError, match="Supported providers"): + await litellm.adecisions( + model="perplexity/pplx-decider-v1-27b", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + api_key="caller-key", + custom_llm_provider="", + ) + + assert len(respx_mock.calls) == 0 + + +@pytest.mark.asyncio +async def test_invalid_question_is_rejected_before_http(respx_mock: respx.MockRouter) -> None: + with pytest.raises(litellm.BadRequestError, match="Invalid Decisions request"): + await litellm.adecisions( + model="perplexity/pplx-decider-v1-27b", + state="review", + questions={"sentiment": {"type": "choice"}}, + api_key="caller-key", + ) + + assert len(respx_mock.calls) == 0 + + +def test_upstream_bad_request_maps_to_litellm_error(respx_mock: respx.MockRouter) -> None: + respx_mock.post("https://api.perplexity.ai/v1/decisions").respond( + status_code=400, + json={"error": {"message": "invalid decision"}}, + ) + + with pytest.raises(litellm.BadRequestError): + litellm.decisions( + model="perplexity/pplx-decider-v1-27b", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + api_key="caller-key", + ) + + +def test_server_key_is_sent_to_an_explicit_api_base( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("PERPLEXITYAI_API_KEY", "server-key") + monkeypatch.delenv("PERPLEXITY_API_KEY", raising=False) + route: Final = respx_mock.post("https://egress.example/perplexity/v1/decisions").respond(json=_RESPONSE) + + litellm.decisions( + model="perplexity/pplx-decider-v1-27b", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + api_base="https://egress.example/perplexity", + ) + + assert route.call_count == 1 + assert route.calls[0].request.headers["authorization"] == "Bearer server-key" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ("cloudflare/clef", "cloudflare/@cf/cloudflare/clef")) +@pytest.mark.parametrize("wrapped", (False, True)) +async def test_cloudflare_clef_resolves_model_and_response_envelope( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, + model: str, + wrapped: bool, +) -> None: + monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct") + monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key") + monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False) + response_body: Final[Mapping[str, object]] = ( + {"result": _RESPONSE, "success": True, "errors": [], "messages": []} if wrapped else _RESPONSE + ) + route: Final = respx_mock.post( + "https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef" + ).respond(json=response_body) + + response: Final = await litellm.adecisions( + model=model, + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + assert route.called + request: Final = respx_mock.calls[0].request + assert request.headers["authorization"] == "Bearer cloudflare-key" + assert json.loads(request.content) == { + "model": "clef", + "state": "review", + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + } + assert response.answers == DecisionsResponse.model_validate(_RESPONSE).answers + assert response._hidden_params["model"] == "cloudflare/@cf/cloudflare/clef" + + +@pytest.mark.asyncio +async def test_cloudflare_clef_flash_uses_flash_endpoint_and_request_model( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct") + monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key") + monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False) + route: Final = respx_mock.post( + "https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef-flash" + ).respond(json=_RESPONSE) + + await litellm.adecisions( + model="cloudflare/clef-flash", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + assert route.called + assert json.loads(respx_mock.calls[0].request.content)["model"] == "clef-flash" + + +@pytest.mark.asyncio +async def test_cloudflare_api_base_from_env_uses_workers_ai_run_path( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("CLOUDFLARE_API_BASE", "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1") + monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key") + monkeypatch.delenv("CLOUDFLARE_ACCOUNT_ID", raising=False) + route: Final = respx_mock.post( + "https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef" + ).respond(json=_RESPONSE) + + await litellm.adecisions( + model="cloudflare/clef", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + assert route.called + + +@pytest.mark.asyncio +async def test_cloudflare_requires_account_id_or_api_base_before_http( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("CLOUDFLARE_ACCOUNT_ID", raising=False) + monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False) + monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key") + + with pytest.raises(litellm.BadRequestError, match="Missing CLOUDFLARE_ACCOUNT_ID - set CLOUDFLARE_ACCOUNT_ID"): + await litellm.adecisions( + model="cloudflare/clef", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + assert len(respx_mock.calls) == 0 + + +@pytest.mark.asyncio +async def test_cloudflare_clef_cost_uses_the_model_cost_map( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct") + monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key") + monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False) + respx_mock.post("https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef").respond( + json=_RESPONSE + ) + + response: Final = await litellm.adecisions( + model="cloudflare/clef", + state="review", + questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + ) + + cost: Final = litellm.completion_cost(completion_response=response) + clef_cost: Final = litellm.model_cost["cloudflare/@cf/cloudflare/clef"] + expected_cost: Final = _INPUT_TOKENS * float(clef_cost["input_cost_per_token"]) + _OUTPUT_TOKENS * float( + clef_cost["output_cost_per_token"] + ) + + assert expected_cost > 0 + assert cost == pytest.approx(expected_cost) + + +@pytest.mark.asyncio +async def test_strands_decider_requires_api_base_before_http( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False) + monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False) + + with pytest.raises(litellm.BadRequestError, match="api_base is required"): + await litellm.adecisions( + model="strands_decider/strands-decider-2B-hobson-v19", + state="review", + questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}}, + ) + + assert len(respx_mock.calls) == 0 + + +@pytest.mark.asyncio +async def test_strands_decider_without_key_preserves_response_extras( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False) + monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False) + route: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE) + + response: Final = await litellm.adecisions( + model="strands_decider/strands-decider-2B-hobson-v19", + state="review", + questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}}, + api_base="https://strands.example", + ) + + assert route.called + assert "authorization" not in respx_mock.calls[0].request.headers + assert response.model_extra["latency_ms"] == _STRANDS_RESPONSE["latency_ms"] + severity: Final = response.answers["severity"] + assert isinstance(severity, ScoreAnswer) + assert severity.legend == {"0": "none", "1": "low", "2": "high"} + + +@pytest.mark.asyncio +async def test_strands_decider_uses_key_from_matching_environment_base( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("STRANDS_DECIDER_API_BASE", "https://strands.example") + monkeypatch.setenv("STRANDS_DECIDER_API_KEY", "strands-key") + route: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE) + + await litellm.adecisions( + model="strands_decider/strands-decider-2B-hobson-v19", + state="review", + questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}}, + api_base="https://strands.example", + ) + + assert route.called + assert respx_mock.calls[0].request.headers["authorization"] == "Bearer strands-key" + + +@pytest.mark.asyncio +async def test_strands_decider_provider_resolution_and_router_dispatch( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False) + monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False) + provider_resolution: Final = litellm.get_llm_provider("strands_decider/strands-decider-2B-hobson-v19") + router: Final = litellm.Router( + model_list=[ + { + "model_name": "strands", + "litellm_params": { + "model": "strands_decider/strands-decider-2B-hobson-v19", + "api_base": "https://strands.example", + }, + } + ] + ) + route: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE) + + response: Final = await router.adecisions( + model="strands", + state="review", + questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}}, + ) + + assert provider_resolution[:2] == ("strands-decider-2B-hobson-v19", "strands_decider") + assert route.called + assert response.model == _STRANDS_RESPONSE["model"] diff --git a/tests/unit/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py index 47c4576f91f..941e44feb26 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -1,5 +1,6 @@ """Test health check helper functions""" +import json import socket import struct import zlib @@ -8,6 +9,7 @@ from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest +import respx import litellm from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME @@ -548,3 +550,100 @@ def test_ocr_health_check_document_raises_without_the_extension(): _ocr_health_check_document(model="mistral/mistral-ocr-latest", custom_llm_provider="mistral") finally: NATIVE_OCR_HEALTH_CHECK_DOCUMENT.reset() + + +@pytest.mark.parametrize( + ("model", "upstream_url"), + ( + ("perplexity/pplx-decider-v1-27b", "https://api.perplexity.ai/v1/decisions"), + ("cloudflare/clef", "https://api.cloudflare.com/client/v4/accounts/acct-1/ai/run/@cf/cloudflare/clef"), + ), +) +async def test_ahealth_check_probes_evaluation_models_through_the_decisions_api( + model: str, + upstream_url: str, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct-1") + monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post(upstream_url).respond( + json={ + "model": model, + "answers": {"reachable": {"type": "noul", "noul": 1.0}}, + "usage": {"input_tokens": 12, "output_tokens": 1}, + } + ) + + result: Final = await ahealth_check({"model": model, "api_key": "sk-test"}, mode=None) + + assert "error" not in result, result + assert upstream.called + sent: Final = json.loads(upstream.calls[0].request.content) + assert sent["state"] == "health check" + assert sent["questions"]["reachable"]["type"] == "noul" + + +@pytest.mark.asyncio +async def test_ahealth_check_evaluation_uses_configured_probe_state_and_questions( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond( + json={ + "model": "perplexity/pplx-decider-v1-27b", + "answers": {"ok": {"type": "noul", "noul": 1.0}}, + "usage": {"input_tokens": 12, "output_tokens": 1}, + } + ) + + result: Final = await ahealth_check( + { + "model": "perplexity/pplx-decider-v1-27b", + "api_key": "sk-test", + "state": "custom probe", + "questions": {"ok": {"type": "noul", "instructions": "Is it ok?"}}, + }, + mode=None, + ) + + assert "error" not in result, result + assert upstream.called + sent: Final = json.loads(upstream.calls[0].request.content) + assert sent["state"] == "custom probe" + assert set(sent["questions"]) == {"ok"} + + +@pytest.mark.asyncio +async def test_ahealth_check_probes_strands_through_decisions_without_mode( + local_model_cost_map: None, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False) + monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post("http://strands.local:8080/v1/systemone").respond( + json={ + "model": "strands-decider-2B-hobson-v19", + "answers": {"reachable": {"type": "noul", "noul": 1.0}}, + "usage": {"input_tokens": 12, "output_tokens": 1}, + } + ) + + result: Final = await ahealth_check( + { + "model": "strands_decider/strands-decider-2B-hobson-v19", + "api_base": "http://strands.local:8080", + }, + mode=None, + ) + + assert "error" not in result, result + assert upstream.called + assert "authorization" not in upstream.calls[0].request.headers diff --git a/tests/unit/llms/azure/test_azure_common_utils.py b/tests/unit/llms/azure/test_azure_common_utils.py index caf941ebd19..5f254a3bc6c 100644 --- a/tests/unit/llms/azure/test_azure_common_utils.py +++ b/tests/unit/llms/azure/test_azure_common_utils.py @@ -470,6 +470,7 @@ def test_default_max_retries_env_var_reaches_azure_sdk_client(): "allm_passthrough_route", "llm_passthrough_route", "asearch", + "adecisions", "avector_store_create", "avector_store_search", "acreate_skill", diff --git a/tests/unit/proxy/decisions_endpoints/__init__.py b/tests/unit/proxy/decisions_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/decisions_endpoints/test_endpoints.py b/tests/unit/proxy/decisions_endpoints/test_endpoints.py new file mode 100644 index 00000000000..6b1ac9e3404 --- /dev/null +++ b/tests/unit/proxy/decisions_endpoints/test_endpoints.py @@ -0,0 +1,326 @@ +from __future__ import annotations + +import asyncio +import json +from collections.abc import AsyncGenerator, Iterator, Mapping +from contextlib import asynccontextmanager +from pathlib import Path +from typing import Final + +import pytest +import respx +from fastapi import FastAPI +from fastapi.routing import APIRoute +from fastapi.testclient import TestClient +from starlette.routing import Match + +import litellm +from litellm.proxy._lazy_features import LAZY_FEATURES, LazyFeature, attach_lazy_features +from litellm.proxy.decisions_endpoints.endpoints import decisions +from litellm.proxy.pass_through_endpoints.pass_through_endpoints import SafeRouteAdder +from litellm.proxy.proxy_server import ( + app, + cleanup_router_config_variables, + initialize, +) + +_INPUT_TOKENS: Final[int] = 367 +_OUTPUT_TOKENS: Final[int] = 3 +_RESPONSE: Final[Mapping[str, object]] = { + "model": "pplx-decider-v1-27b", + "answers": { + "is_defect": {"type": "noul", "noul": 0.9}, + }, + "usage": {"input_tokens": _INPUT_TOKENS, "output_tokens": _OUTPUT_TOKENS}, +} +_STRANDS_RESPONSE: Final[Mapping[str, object]] = { + "model": "strands-decider-2B-hobson-v19", + "answers": { + "is_defect": {"type": "noul", "noul": 0.9}, + }, + "usage": {"input_tokens": 216, "output_tokens": 3}, + "latency_ms": 3722.17, +} +_REQUEST: Final[Mapping[str, object]] = { + "model": "decider", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, +} + + +@pytest.fixture +def client(monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: + monkeypatch.setenv("OPENAI_API_KEY", "fake-openai-key") + monkeypatch.setenv("OPENAI_API_BASE", "https://fake-openai.example") + monkeypatch.setenv("REDIS_HOST", "localhost") + cleanup_router_config_variables() + config_path: Final = Path(__file__).parents[1] / "test_configs" / "test_config_no_auth.yaml" + asyncio.run(initialize(config=str(config_path), debug=True)) + monkeypatch.setattr( + litellm.proxy.proxy_server, + "llm_router", + litellm.Router( + model_list=[ + { + "model_name": "decider", + "litellm_params": { + "model": "perplexity/pplx-decider-v1-27b", + "api_key": "test-key", + }, + } + ] + ), + ) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield TestClient(app) + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.parametrize("endpoint", ("/v1/decisions", "/decisions")) +def test_proxy_decisions_route_returns_answers_and_cost( + client: TestClient, + respx_mock: respx.MockRouter, + endpoint: str, +) -> None: + upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + + response: Final = client.post(endpoint, json=_REQUEST) + + assert response.status_code == 200, response.text + assert response.json()["answers"] == _RESPONSE["answers"] + assert "_hidden_params" not in response.json() + perplexity_cost: Final = litellm.model_cost["perplexity/pplx-decider-v1-27b"] + expected_cost: Final = _INPUT_TOKENS * float(perplexity_cost["input_cost_per_token"]) + _OUTPUT_TOKENS * float( + perplexity_cost["output_cost_per_token"] + ) + + assert expected_cost > 0 + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected_cost) + assert upstream.called + assert json.loads(upstream.calls[0].request.content) == { + "model": "pplx-decider-v1-27b", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + } + assert upstream.calls[0].request.headers["authorization"] == "Bearer test-key" + + +def test_proxy_decisions_dispatches_typesafe_deployment( + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + router: Final = litellm.Router( + model_list=[ + { + "model_name": "jev", + "litellm_params": { + "model": "typesafe/jev-latest", + "api_key": "k", + }, + } + ] + ) + monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) + upstream: Final = respx_mock.post("https://api.typesafe.ai/v1/systemone").respond(json=_RESPONSE) + + response: Final = client.post( + "/v1/decisions", + json={ + "model": "jev", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + }, + ) + + assert response.status_code == 200, response.text + assert response.json()["answers"] == _RESPONSE["answers"] + assert upstream.called + assert json.loads(upstream.calls[0].request.content) == { + "model": "jev-latest", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + } + + +def test_proxy_decisions_sends_the_env_key_to_the_deployment_api_base( + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("PERPLEXITYAI_API_KEY", "server-key") + monkeypatch.delenv("PERPLEXITY_API_KEY", raising=False) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "decider", + "litellm_params": { + "model": "perplexity/pplx-decider-v1-27b", + "api_base": "https://egress.example/perplexity", + }, + } + ] + ) + monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) + upstream: Final = respx_mock.post("https://egress.example/perplexity/v1/decisions").respond(json=_RESPONSE) + + response: Final = client.post("/v1/decisions", json=_REQUEST) + + assert response.status_code == 200, response.text + assert upstream.call_count == 1 + assert upstream.calls[0].request.headers["authorization"] == "Bearer server-key" + + +def test_proxy_decisions_unknown_model_is_a_client_error( + client: TestClient, + respx_mock: respx.MockRouter, +) -> None: + response: Final = client.post( + "/v1/decisions", + json={ + "model": "missing-model", + "state": "review", + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + }, + ) + + assert 400 <= response.status_code < 500, response.text + assert len(respx_mock.calls) == 0 + + +@pytest.mark.parametrize( + "request_body", + ( + { + "model": "decider", + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + }, + { + "model": "decider", + "state": {"source": "proxy-test"}, + }, + ), + ids=("missing_state", "missing_questions"), +) +def test_proxy_decisions_missing_required_field_is_a_client_error( + client: TestClient, + respx_mock: respx.MockRouter, + request_body: Mapping[str, object], +) -> None: + upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + + response: Final = client.post("/v1/decisions", json=request_body) + + assert response.status_code == 400, response.text + assert not upstream.called + + +def test_proxy_decisions_dispatches_strands_decider( + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False) + monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "strands", + "litellm_params": { + "model": "strands_decider/strands-decider-2B-hobson-v19", + "api_base": "https://strands.example", + }, + } + ] + ) + monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) + upstream: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE) + + response: Final = client.post( + "/v1/decisions", + json={ + "model": "strands", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + }, + ) + + assert response.status_code == 200, response.text + assert response.json()["answers"] == _STRANDS_RESPONSE["answers"] + assert upstream.called + assert json.loads(upstream.calls[0].request.content) == { + "model": "strands-decider-2B-hobson-v19", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + } + assert "authorization" not in upstream.calls[0].request.headers + + +def test_proxy_decisions_without_model_uses_the_proxy_default_model( + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setattr(litellm.proxy.proxy_server, "user_model", "decider") + upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + + response: Final = client.post( + "/v1/decisions", json={key: value for key, value in _REQUEST.items() if key != "model"} + ) + + assert response.status_code == 200, response.text + assert response.json()["answers"] == _RESPONSE["answers"] + assert upstream.called + assert json.loads(upstream.calls[0].request.content)["model"] == "pplx-decider-v1-27b" + + +def _decisions_feature() -> LazyFeature: + return next(feature for feature in LAZY_FEATURES if feature.name == "decisions") + + +def _serving_endpoint(bare: FastAPI, path: str) -> object: + scope: Final = {"type": "http", "method": "POST", "path": path, "root_path": "", "query_string": b"", "headers": ()} + return next( + route.endpoint for route in bare.routes if isinstance(route, APIRoute) and route.matches(scope)[0] is Match.FULL + ) + + +def test_a_config_pass_through_at_v1_decisions_keeps_its_route_and_the_native_api_serves_decisions( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("LITELLM_DISABLE_LAZY_ROUTES", raising=False) + + async def pass_through() -> dict[str, str]: + return {"served_by": "pass-through"} + + bare: Final = FastAPI() + attach_lazy_features(bare, (_decisions_feature(),)) + SafeRouteAdder.add_api_route_if_not_exists(bare, "/v1/decisions", pass_through, ["POST"]) + with TestClient(bare) as client: + assert client.post("/v1/decisions", json={"model": "gpt-6-luna"}).json() == {"served_by": "pass-through"} + assert _serving_endpoint(bare, "/v1/decisions") is pass_through + assert _serving_endpoint(bare, "/decisions") is decisions + + +def test_with_lazy_routes_disabled_a_config_pass_through_at_v1_decisions_still_wins( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("LITELLM_DISABLE_LAZY_ROUTES", "true") + + async def pass_through() -> dict[str, str]: + return {"served_by": "pass-through"} + + @asynccontextmanager + async def loads_the_config(app_: FastAPI) -> AsyncGenerator[None]: + assert SafeRouteAdder.add_api_route_if_not_exists(app_, "/v1/decisions", pass_through, ["POST"]), ( + "the native route registered at startup must not block the config pass-through" + ) + yield + + bare: Final = FastAPI(lifespan=loads_the_config) + attach_lazy_features(bare, (_decisions_feature(),)) + with TestClient(bare) as client: + assert client.post("/v1/decisions", json={"model": "gpt-6-luna"}).json() == {"served_by": "pass-through"} + assert _serving_endpoint(bare, "/v1/decisions") is pass_through + assert _serving_endpoint(bare, "/decisions") is decisions diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index c6c81c14b16..31321254d94 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -9,15 +9,16 @@ from collections.abc import Callable, Mapping from contextlib import ExitStack, contextmanager from dataclasses import dataclass from io import BytesIO -from types import MappingProxyType, SimpleNamespace +from types import MappingProxyType, ModuleType, SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest import respx -from fastapi import HTTPException, Request, Response, UploadFile +from fastapi import APIRouter, FastAPI, HTTPException, Request, Response, UploadFile from fastapi.responses import StreamingResponse +from fastapi.testclient import TestClient from pydantic import TypeAdapter, ValidationError from starlette.datastructures import FormData, Headers, QueryParams from starlette.datastructures import UploadFile as StarletteUploadFile @@ -27,12 +28,14 @@ from litellm._logging import verbose_proxy_logger from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.proxy._lazy_features import LazyFeature, attach_lazy_features from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, HttpPassThroughEndpointHelpers, InitPassThroughEndpointHelpers, + SafeRouteAdder, _registered_pass_through_routes, _truncate_upstream_error_body, _with_trace_context, @@ -8319,3 +8322,31 @@ async def test_pass_throughs_stay_open_while_a_db_sync_reads_the_database(tmp_pa await sync assert (config_during_sync.status_code, db_during_sync.status_code) == (200, 200) + + +def _lazy_feature(monkeypatch: pytest.MonkeyPatch, name: str, path: str) -> LazyFeature: + async def served() -> dict[str, str]: + return {"served_by": name} + + router: Final = APIRouter() + router.add_api_route(path, served, methods=["POST"]) + module: Final = ModuleType(f"tests.unit.proxy.pass_through_endpoints.lazy_fixture_{name}") + module.router = router # pyright: ignore[reportAttributeAccessIssue] # fixture module built at test time + monkeypatch.setitem(sys.modules, module.__name__, module) + return LazyFeature(name=name, module_path=module.__name__, path_prefixes=(path,)) + + +def test_a_pass_through_added_after_a_lazy_feature_loaded_takes_over_its_path(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("LITELLM_DISABLE_LAZY_ROUTES", raising=False) + + async def pass_through() -> dict[str, str]: + return {"served_by": "pass-through"} + + app: Final = FastAPI() + attach_lazy_features(app, (_lazy_feature(monkeypatch, "decider", "/v1/decider"),)) + with TestClient(app) as client: + assert client.post("/v1/decider").json() == {"served_by": "decider"} + assert SafeRouteAdder.add_api_route_if_not_exists(app, "/v1/decider", pass_through, ["POST"]) + assert client.post("/v1/decider").json() == {"served_by": "pass-through"} + assert not SafeRouteAdder.add_api_route_if_not_exists(app, "/v1/decider", pass_through, ["POST"]) + assert client.post("/v1/decider").json() == {"served_by": "pass-through"} diff --git a/tests/unit/proxy/public_endpoints/test_public_endpoints.py b/tests/unit/proxy/public_endpoints/test_public_endpoints.py index be309a67d58..448ad9c712e 100644 --- a/tests/unit/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/unit/proxy/public_endpoints/test_public_endpoints.py @@ -436,10 +436,12 @@ ADD_MODEL_UNLISTED_PROVIDERS: Final = frozenset( "sagemaker_nova", "scaleway", "stability", + "strands_decider", "synthetic", "tensormesh", "text-completion-inception", "transcribe", + "typesafe", "valkey", "xiaomi_mimo", "zai", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index ba7dd127aae..2c14143cfe7 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -4586,6 +4586,23 @@ export interface paths { patch?: never; trace?: never; }; + "/decisions": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Decisions */ + post: operations["decisions_decisions_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/delete/allowed_ip": { parameters: { query?: never; @@ -19790,6 +19807,23 @@ export interface paths { patch?: never; trace?: never; }; + "/v1/decisions": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Decisions */ + post: operations["decisions_v1_decisions_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/v1/embeddings": { parameters: { query?: never; @@ -28092,7 +28126,7 @@ export interface components { * CallTypes * @enum {string} */ - CallTypes: "embedding" | "aembedding" | "completion" | "acompletion" | "atext_completion" | "text_completion" | "image_generation" | "aimage_generation" | "image_edit" | "aimage_edit" | "moderation" | "amoderation" | "atranscription" | "transcription" | "aspeech" | "speech" | "rerank" | "arerank" | "search" | "asearch" | "_arealtime" | "_aresponses_websocket" | "create_batch" | "acreate_batch" | "aretrieve_batch" | "retrieve_batch" | "acancel_batch" | "cancel_batch" | "pass_through_endpoint" | "anthropic_messages" | "aanthropic_messages" | "get_assistants" | "aget_assistants" | "create_assistants" | "acreate_assistants" | "delete_assistant" | "adelete_assistant" | "acreate_thread" | "create_thread" | "aget_thread" | "get_thread" | "a_add_message" | "add_message" | "aget_messages" | "get_messages" | "arun_thread" | "run_thread" | "arun_thread_stream" | "run_thread_stream" | "afile_retrieve" | "file_retrieve" | "afile_delete" | "file_delete" | "afile_list" | "file_list" | "acreate_file" | "create_file" | "afile_content" | "file_content" | "create_fine_tuning_job" | "acreate_fine_tuning_job" | "create_video" | "acreate_video" | "video_generation" | "avideo_generation" | "avideo_retrieve" | "video_retrieve" | "avideo_content" | "video_content" | "video_remix" | "avideo_remix" | "video_list" | "avideo_list" | "video_retrieve_job" | "avideo_retrieve_job" | "video_delete" | "avideo_delete" | "video_create_character" | "avideo_create_character" | "video_get_character" | "avideo_get_character" | "video_edit" | "avideo_edit" | "video_extension" | "avideo_extension" | "vector_store_file_create" | "avector_store_file_create" | "vector_store_file_list" | "avector_store_file_list" | "vector_store_file_retrieve" | "avector_store_file_retrieve" | "vector_store_file_content" | "avector_store_file_content" | "vector_store_file_update" | "avector_store_file_update" | "vector_store_file_delete" | "avector_store_file_delete" | "vector_store_create" | "avector_store_create" | "vector_store_search" | "avector_store_search" | "ingest" | "aingest" | "query" | "aquery" | "create_interaction" | "acreate_interaction" | "create_container" | "acreate_container" | "list_containers" | "alist_containers" | "retrieve_container" | "aretrieve_container" | "delete_container" | "adelete_container" | "list_container_files" | "alist_container_files" | "upload_container_file" | "aupload_container_file" | "create_sandbox" | "acreate_sandbox" | "delete_sandbox" | "adelete_sandbox" | "run_code" | "arun_code" | "code_interpreter_tool" | "acode_interpreter_tool" | "acancel_fine_tuning_job" | "cancel_fine_tuning_job" | "alist_fine_tuning_jobs" | "list_fine_tuning_jobs" | "aretrieve_fine_tuning_job" | "retrieve_fine_tuning_job" | "responses" | "aresponses" | "alist_input_items" | "llm_passthrough_route" | "allm_passthrough_route" | "generate_content" | "agenerate_content" | "generate_content_stream" | "agenerate_content_stream" | "ocr" | "aocr" | "call_mcp_tool" | "list_mcp_tools" | "asend_message" | "send_message" | "acreate_skill"; + CallTypes: "embedding" | "aembedding" | "completion" | "acompletion" | "atext_completion" | "text_completion" | "image_generation" | "aimage_generation" | "image_edit" | "aimage_edit" | "moderation" | "amoderation" | "atranscription" | "transcription" | "aspeech" | "speech" | "rerank" | "arerank" | "search" | "asearch" | "decisions" | "adecisions" | "_arealtime" | "_aresponses_websocket" | "create_batch" | "acreate_batch" | "aretrieve_batch" | "retrieve_batch" | "acancel_batch" | "cancel_batch" | "pass_through_endpoint" | "anthropic_messages" | "aanthropic_messages" | "get_assistants" | "aget_assistants" | "create_assistants" | "acreate_assistants" | "delete_assistant" | "adelete_assistant" | "acreate_thread" | "create_thread" | "aget_thread" | "get_thread" | "a_add_message" | "add_message" | "aget_messages" | "get_messages" | "arun_thread" | "run_thread" | "arun_thread_stream" | "run_thread_stream" | "afile_retrieve" | "file_retrieve" | "afile_delete" | "file_delete" | "afile_list" | "file_list" | "acreate_file" | "create_file" | "afile_content" | "file_content" | "create_fine_tuning_job" | "acreate_fine_tuning_job" | "create_video" | "acreate_video" | "video_generation" | "avideo_generation" | "avideo_retrieve" | "video_retrieve" | "avideo_content" | "video_content" | "video_remix" | "avideo_remix" | "video_list" | "avideo_list" | "video_retrieve_job" | "avideo_retrieve_job" | "video_delete" | "avideo_delete" | "video_create_character" | "avideo_create_character" | "video_get_character" | "avideo_get_character" | "video_edit" | "avideo_edit" | "video_extension" | "avideo_extension" | "vector_store_file_create" | "avector_store_file_create" | "vector_store_file_list" | "avector_store_file_list" | "vector_store_file_retrieve" | "avector_store_file_retrieve" | "vector_store_file_content" | "avector_store_file_content" | "vector_store_file_update" | "avector_store_file_update" | "vector_store_file_delete" | "avector_store_file_delete" | "vector_store_create" | "avector_store_create" | "vector_store_search" | "avector_store_search" | "ingest" | "aingest" | "query" | "aquery" | "create_interaction" | "acreate_interaction" | "create_container" | "acreate_container" | "list_containers" | "alist_containers" | "retrieve_container" | "aretrieve_container" | "delete_container" | "adelete_container" | "list_container_files" | "alist_container_files" | "upload_container_file" | "aupload_container_file" | "create_sandbox" | "acreate_sandbox" | "delete_sandbox" | "adelete_sandbox" | "run_code" | "arun_code" | "code_interpreter_tool" | "acode_interpreter_tool" | "acancel_fine_tuning_job" | "cancel_fine_tuning_job" | "alist_fine_tuning_jobs" | "list_fine_tuning_jobs" | "aretrieve_fine_tuning_job" | "retrieve_fine_tuning_job" | "responses" | "aresponses" | "alist_input_items" | "llm_passthrough_route" | "allm_passthrough_route" | "generate_content" | "agenerate_content" | "generate_content_stream" | "agenerate_content_stream" | "ocr" | "aocr" | "call_mcp_tool" | "list_mcp_tools" | "asend_message" | "send_message" | "acreate_skill"; /** CallbackDelete */ CallbackDelete: { /** Callback Name */ @@ -56843,6 +56877,26 @@ export interface operations { }; }; }; + decisions_decisions_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; delete_allowed_ip_delete_allowed_ip_post: { parameters: { query?: never; @@ -76518,6 +76572,26 @@ export interface operations { }; }; }; + decisions_v1_decisions_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; embeddings_v1_embeddings_post: { parameters: { query?: never; From ad8babae33375948646b0a2c63d56be5cf590ced Mon Sep 17 00:00:00 2001 From: moe-berri Date: Sat, 3 Oct 2026 10:38:40 -0700 Subject: [PATCH 04/82] fix(lens): paginate trace reads within ClickHouse limits (#44384) --- .../crates/python-bridge/src/routes/traces.rs | 43 +- .../crates/storage-clickhouse/src/read.rs | 7 + .../storage-clickhouse/tests/transport.rs | 33 ++ .../traces-clickhouse/query/spend_batch.sql | 27 ++ .../query/trace_span_batch.sql | 30 ++ .../crates/traces-clickhouse/src/error.rs | 4 + .../crates/traces-clickhouse/src/lib.rs | 3 +- .../crates/traces-clickhouse/src/reads.rs | 249 +++++++--- .../traces-clickhouse/src/span_batches.rs | 181 ++++++++ .../crates/traces-clickhouse/tests/reads.rs | 436 ++++++++++++++++++ .../crates/traces/src/resolve/view.rs | 3 + litellm-rust/crates/traces/src/view.rs | 12 +- litellm/integrations/otel/model/spans.py | 1 + litellm/proxy/tracing_endpoints.py | 34 +- litellm/rust_bridge/_native.pyi | 4 +- litellm/rust_bridge/trace/generated/types.py | 2 + litellm/rust_bridge/trace/storage.py | 15 +- litellm/tracing/receiver.py | 11 +- .../trace_codegen/schemas/traces/Trace.json | 15 +- .../schemas/traces/TracePage.json | 5 + tests/test_litellm/tracing/test_receiver.py | 7 +- tests/unit/proxy/test_tracing_endpoints.py | 72 ++- .../src/components/networking.tsx | 9 +- .../AgentTracesSection.integration.test.tsx | 19 + .../TraceView/AgentTracesSection.tsx | 2 + .../view_logs/TraceView/AgentTracesTable.tsx | 27 +- .../view_logs/TraceView/TraceDrawer.test.tsx | 37 +- .../view_logs/TraceView/TraceDrawer.tsx | 64 ++- .../view_logs/TraceView/tracesApi.ts | 4 +- .../view_logs/TraceView/useAgentTraces.ts | 26 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 6 + 31 files changed, 1258 insertions(+), 130 deletions(-) create mode 100644 litellm-rust/crates/traces-clickhouse/query/spend_batch.sql create mode 100644 litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql create mode 100644 litellm-rust/crates/traces-clickhouse/src/span_batches.rs create mode 100644 litellm-rust/crates/traces-clickhouse/tests/reads.rs diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 644627b05fd..77a98b3980a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -35,13 +35,14 @@ fn map_error_ref(error: &Error) -> PyErr { use litellm_storage_clickhouse::Error as StorageError; match error { - Error::Decode(litellm_traces::Error::TooLarge) | Error::InsertTooLarge => { - PyOverflowError::new_err(error.to_string()) - } + Error::Decode(litellm_traces::Error::TooLarge) + | Error::InsertTooLarge + | Error::ReadTooLarge => PyOverflowError::new_err(error.to_string()), Error::InvalidRow | Error::InvalidTable | Error::InvalidCursor(_) | Error::AmbiguousTrace + | Error::TraceChanged | Error::Decode(_) | Error::InvalidSchema | Error::InvalidQuery @@ -233,26 +234,44 @@ impl NativeTraceStorage { ) } + #[pyo3(signature = (trace_id, scope, trace_ref, cursor=None, page_size=None))] fn get_trace<'py>( &self, py: Python<'py>, trace_id: String, #[pyo3(from_py_with = litellm_host_python::from_py_argument)] scope: ReadAccessParams, trace_ref: String, + cursor: Option, + page_size: Option, ) -> PyResult> { let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; let connection = self.config.storage().reader().clone(); crate::execution::run_async( py, async move { - litellm_traces_clickhouse::get_trace( - &client, - &connection, - &scope, - &trace_id, - &trace_ref, - ) - .await + if let Some(page_size) = page_size { + litellm_traces_clickhouse::get_trace_page( + &client, + &connection, + &scope, + &trace_id, + &trace_ref, + cursor.as_deref(), + page_size, + ) + .await + } else if cursor.is_some() { + Err(Error::InvalidParameters) + } else { + litellm_traces_clickhouse::get_trace( + &client, + &connection, + &scope, + &trace_id, + &trace_ref, + ) + .await + } }, map_error, ) @@ -465,6 +484,8 @@ mod tests { #[case::invalid_export(Error::Decode(litellm_traces::Error::InvalidPayload), "ValueError")] #[case::cursor(Error::InvalidCursor("trace"), "ValueError")] #[case::ambiguous(Error::AmbiguousTrace, "ValueError")] + #[case::changed_snapshot(Error::TraceChanged, "ValueError")] + #[case::read_budget(Error::ReadTooLarge, "OverflowError")] fn trace_read_and_ingest_failures_preserve_public_exception_types( #[case] error: Error, #[case] exception_name: &str, diff --git a/litellm-rust/crates/storage-clickhouse/src/read.rs b/litellm-rust/crates/storage-clickhouse/src/read.rs index 59d5b0de558..c4bfdef393a 100644 --- a/litellm-rust/crates/storage-clickhouse/src/read.rs +++ b/litellm-rust/crates/storage-clickhouse/src/read.rs @@ -110,6 +110,13 @@ pub async fn execute_read( .body(sql.to_owned()); let mut response = request.send().await.map_err(|_| Error::Transport)?; if !response.status().is_success() { + if response + .headers() + .get("x-clickhouse-exception-code") + .is_some_and(|code| code == "396") + { + return Err(Error::ResponseTooLarge); + } return Err(Error::QueryFailed(response.status().as_u16())); } diff --git a/litellm-rust/crates/storage-clickhouse/tests/transport.rs b/litellm-rust/crates/storage-clickhouse/tests/transport.rs index f58e52b4940..1dab575e21b 100644 --- a/litellm-rust/crates/storage-clickhouse/tests/transport.rs +++ b/litellm-rust/crates/storage-clickhouse/tests/transport.rs @@ -111,3 +111,36 @@ async fn typed_fetch_encodes_parameters_and_validates_rows( assert!(matches!(envelope, Err(Error::InvalidResponse))); } } + +#[rstest] +#[case::result_limit("396", true)] +#[case::memory_limit("241", false)] +#[case::timeout("159", false)] +#[case::unknown("", false)] +#[tokio::test] +async fn server_result_limits_allow_smaller_pages_without_retrying_other_failures( + #[case] code: &str, + #[case] result_limit: bool, +) { + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(500).insert_header("X-ClickHouse-Exception-Code", code)) + .expect(1) + .mount(&server) + .await; + let connection = Connection::parse(&server.uri()).unwrap(); + let error = execute_read( + &Client::no_redirect_for_test(), + &connection, + "SELECT 1", + &BTreeMap::new(), + ) + .await + .unwrap_err(); + if result_limit { + assert!(matches!(error, Error::ResponseTooLarge)); + } else { + assert!(matches!(error, Error::QueryFailed(500))); + } +} diff --git a/litellm-rust/crates/traces-clickhouse/query/spend_batch.sql b/litellm-rust/crates/traces-clickhouse/query/spend_batch.sql new file mode 100644 index 00000000000..658169dbe34 --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/query/spend_batch.sql @@ -0,0 +1,27 @@ +SELECT * FROM ( +SELECT request_id, response_id, upstream_response_id, trace_id, span_id, team_id, api_key, user, spend, + toUnixTimestamp64Milli(start_time) AS start_ms +FROM ( + SELECT *, + -- A chat request served through the Responses API returns the upstream `resp_` id to the + -- client but logs LiteLLM's managed `resp_` id, which embeds it. + if(startsWith(response_id, 'resp_'), + extract(tryBase64Decode(substring(response_id, 6)), 'response_id:([^;]+)'), + '') AS upstream_response_id + FROM spend_logs FINAL + WHERE start_time >= fromUnixTimestamp64Milli({start_ms:Int64}) + AND start_time < fromUnixTimestamp64Milli({end_ms:Int64}) + AND ({all_teams:UInt8} = 1 + OR ({user_id:String} != '' AND user = {user_id:String}) + OR has({team_ids:Array(String)}, team_id)) +) +WHERE response_id IN {response_ids:Array(String)} + OR upstream_response_id IN {response_ids:Array(String)} + OR request_id IN {request_ids:Array(String)} + OR (trace_id != '' AND trace_id IN {trace_ids:Array(String)}) +ORDER BY start_time DESC +) +WHERE {has_cursor:UInt8} = 0 + OR (team_id, start_ms, request_id) > ({after_team:String}, {after_ms:Int64}, {after_id:String}) +ORDER BY team_id, start_ms, request_id +LIMIT {page_size:UInt32} diff --git a/litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql b/litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql new file mode 100644 index 00000000000..967ef2fcf1f --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql @@ -0,0 +1,30 @@ +SELECT * FROM ( +SELECT o.TraceId AS trace_id, o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name, + o.ObservationType AS type, toUInt8(o.WrapperCandidate) AS wrapper_candidate, o.AgentName AS agent, + o.Framework AS framework, o.StatusCode AS status, + substringUTF8(o.StatusMessage, 1, 128) AS status_message, + lengthUTF8(o.StatusMessage) > 128 AS error_truncated, + toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns, + o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model, + o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens, + o.LiteLLMRequestId AS litellm_request_id, + o.CallKeys AS call_keys, o.CallEvidence AS call_evidence, + -- Rows written before ToolCallId keep the call id only in their attributes. + if(o.ToolCallId != '' OR o.ObservationType != 'tool', o.ToolCallId, + coalesce(nullIf(o.SpanAttributes['gen_ai.tool.call.id'], ''), nullIf(o.SpanAttributes['tool.id'], ''), '')) + AS tool_call_id, + o.UserId AS user_id, o.TeamId AS team_id, o.ApiKeyHash AS api_key_hash +FROM otel_traces AS o +WHERE o.TraceId = {trace_id:String} + AND ({all_teams:UInt8} = 1 + OR ({user_id:String} != '' AND o.UserId = {user_id:String}) + OR has({team_ids:Array(String)}, o.TeamId)) + AND ({trace_ref:String} = '' OR + hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) = {trace_ref:String}) + AND o.EngineReceivedMs <= {snapshot_ms:UInt64} +ORDER BY o.Timestamp, o.EngineReceivedMs, o.StatusMessage +LIMIT 1 BY o.SpanId +) +WHERE span_id > {after_span_id:String} +ORDER BY span_id +LIMIT {page_size:UInt32} diff --git a/litellm-rust/crates/traces-clickhouse/src/error.rs b/litellm-rust/crates/traces-clickhouse/src/error.rs index ae39fce25b2..fdbc3029a5d 100644 --- a/litellm-rust/crates/traces-clickhouse/src/error.rs +++ b/litellm-rust/crates/traces-clickhouse/src/error.rs @@ -14,6 +14,8 @@ pub enum Error { InvalidResponse, #[error("ClickHouse insert exceeds the encoded size limit")] InsertTooLarge, + #[error("Trace exceeds the interactive read budget; use a filtered trace query")] + ReadTooLarge, #[error("ClickHouse schema setup failed with HTTP status {0}")] SchemaFailed(u16), #[error("ClickHouse schema setup transport failed")] @@ -34,6 +36,8 @@ pub enum Error { InvalidCursor(&'static str), #[error("Multiple traces have this ID; provide trace_ref")] AmbiguousTrace, + #[error("Trace changed while paging; refresh the trace to continue")] + TraceChanged, #[error(transparent)] Decode(#[from] litellm_traces::Error), #[error("trace ingestion task failed")] diff --git a/litellm-rust/crates/traces-clickhouse/src/lib.rs b/litellm-rust/crates/traces-clickhouse/src/lib.rs index fc67df4eba9..d83708d27f1 100644 --- a/litellm-rust/crates/traces-clickhouse/src/lib.rs +++ b/litellm-rust/crates/traces-clickhouse/src/lib.rs @@ -17,6 +17,7 @@ pub mod query; mod query_access; mod reads; mod schema; +mod span_batches; mod span_row; mod sql; mod table; @@ -30,7 +31,7 @@ pub use litellm_storage_clickhouse::{Connection, Parameter}; pub use litellm_traces::{QueryScope, ReadQuery}; pub use query::{QueryHelp, execute_read, query_help, query_sql}; pub use query_access::QueryReaders; -pub use reads::{get_span, get_span_error, get_trace, list_traces}; +pub use reads::{get_span, get_span_error, get_trace, get_trace_page, list_traces}; pub use schema::{ NORMALIZED_FIELD_DEFINITIONS, NormalizedFieldDefinition, ensure_schema, schema_statements, }; diff --git a/litellm-rust/crates/traces-clickhouse/src/reads.rs b/litellm-rust/crates/traces-clickhouse/src/reads.rs index 68c44efbde0..5949b546335 100644 --- a/litellm-rust/crates/traces-clickhouse/src/reads.rs +++ b/litellm-rust/crates/traces-clickhouse/src/reads.rs @@ -1,26 +1,54 @@ //! Scoped trace reads: the trace list, one trace resolved with its spend, and span payloads. -use std::collections::HashMap; +use std::sync::{Arc, LazyLock}; +use std::time::Duration; use base64::{Engine, engine::general_purpose::URL_SAFE}; use litellm_http::Client; -use litellm_storage_clickhouse::fetch; +use litellm_storage_clickhouse::{Query, fetch}; use litellm_traces::{ SpanDetail, SpanErrorPage, SpendLookup, Trace, TracePage, listed_summary, query::named as contracts, resolve_trace, to_ui_content, }; +use moka::future::Cache; use serde::{Deserialize, Serialize}; +use sha2::{Digest, Sha256}; use crate::{ Connection, Error, query::named::{ - ListTraces, ListTracesParams, ReadAccessParams, SpanDetail as SpanDetailQuery, - SpanDetailParams, SpanError, SpanErrorParams, SpendByResponseIds, SpendByResponseIdsParams, - TraceIdentity, TraceIdentityParams, TracePageSpans, TracePageSpansParams, TraceSpans, - TraceSpansParams, + ListTracesParams, ListTracesRow, ReadAccessParams, SpanDetail as SpanDetailQuery, + SpanDetailParams, SpanError, SpanErrorParams, SpendByResponseIdsParams, TraceIdentity, + TraceIdentityParams, TraceSpansParams, }, }; +struct RunCandidates; + +impl Query for RunCandidates { + type Params = ListTracesParams; + type Row = ListTracesRow; + const SQL: &'static str = concat!( + "SELECT * EXCEPT (request_ids), [] AS request_ids FROM (", + include_str!("../query/list_traces.sql"), + ") ORDER BY start_ms DESC, trace_ref DESC" + ); +} + +// Cursor pages share a bounded snapshot so advancing does not resolve the whole graph again. +static TRACE_SNAPSHOTS: LazyLock>> = LazyLock::new(|| { + Cache::builder() + .max_capacity(64 * 1024 * 1024) + .weigher(|_: &String, trace: &Arc| { + serde_json::to_vec(trace.as_ref()) + .ok() + .and_then(|bytes| u32::try_from(bytes.len().saturating_mul(2)).ok()) + .unwrap_or(u32::MAX) + }) + .time_to_live(Duration::from_secs(120)) + .build() +}); + const NANOS_PER_MS: i64 = 1_000_000; const SPEND_WINDOW_MS: i64 = 30 * 60 * 1000; @@ -121,8 +149,8 @@ async fn spend( start_ms: start_ns.div_euclid(NANOS_PER_MS) - SPEND_WINDOW_MS, end_ms: end_ns.div_euclid(NANOS_PER_MS) + SPEND_WINDOW_MS, }); - match fetch::(client, connection, ¶ms).await { - Ok(rows) => rows.into_iter().map(|row| row.0).collect(), + match crate::span_batches::read_spend(client, connection, params).await { + Ok(rows) => rows, Err(error) => { tracing::warn!(%error, "trace spend lookup unavailable"); Vec::new() @@ -139,71 +167,43 @@ pub async fn list_traces( cursor: Option<&str>, limit: u32, ) -> Result { + if limit == 0 { + return Err(Error::InvalidParameters); + } let (cursor_ms, cursor_trace_id) = trace_position(cursor)?; - let params = ListTracesParams::from(contracts::ListTracesParams { + let mut params = ListTracesParams::from(contracts::ListTracesParams { access: access.clone(), start_ms, end_ms, cursor_ms, cursor_trace_id, - limit, + limit: limit.min(500), }); - let page: Vec = fetch::(client, connection, ¶ms) - .await? - .into_iter() - .map(|row| row.0) - .collect(); + let page: Vec = loop { + match fetch::(client, connection, ¶ms).await { + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) if params.0.limit > 1 => { + params.0.limit /= 2; + } + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { + return Err(Error::ReadTooLarge); + } + result => break result?.into_iter().map(|row| row.0).collect(), + } + }; let next_cursor = page .last() - .filter(|_| page.len() == limit as usize) + .filter(|_| page.len() == params.0.limit as usize) .map(|last| encode_cursor(&(last.start_ms, &last.trace_ref))); - let (Some(page_start), Some(page_end)) = ( - page.iter().map(|row| row.start_ms).min(), - page.iter().map(|row| row.start_ms + row.duration_ms).max(), - ) else { - return Ok(TracePage { - data: Vec::new(), - next_cursor, - }); - }; - let span_params = TracePageSpansParams::from(contracts::TracePageSpansParams { - access: access.clone(), - trace_refs: page.iter().map(|row| row.trace_ref.clone()).collect(), - start_ms: page_start, - end_ms: page_end + 1, - }); - let span_rows: Vec = - fetch::(client, connection, &span_params) - .await? - .into_iter() - .map(|row| row.0) - .collect(); - let spend_rows = spend(client, connection, access, &span_rows).await; - let mut by_trace: HashMap<(String, String, String), Vec> = - HashMap::new(); - for span in span_rows { - let key = ( - span.team_id.clone(), - span.api_key_hash.clone(), - span.trace_id.clone(), - ); - by_trace.entry(key).or_default().push(span); + let mut data = Vec::with_capacity(page.len()); + for row in &page { + let summary = + match get_trace(client, connection, access, &row.trace_id, &row.trace_ref).await { + Ok(trace) => trace.map_or_else(|| listed_summary(row), |trace| trace.summary), + Err(Error::ReadTooLarge) => listed_summary(row), + Err(error) => return Err(error), + }; + data.push(summary); } - let data = page - .iter() - .map(|row| { - let spans = by_trace - .get(&( - row.team_id.clone(), - row.api_key_hash.clone(), - row.trace_id.clone(), - )) - .map(Vec::as_slice) - .unwrap_or_default(); - resolve_trace(&row.trace_id, &row.trace_ref, spans, &spend_rows) - .map_or_else(|| listed_summary(row), |trace| trace.summary) - }) - .collect(); Ok(TracePage { data, next_cursor }) } @@ -222,11 +222,7 @@ pub async fn get_trace( trace_id: trace_id.to_owned(), trace_ref: trace_ref.clone(), }; - let rows: Vec = fetch::(client, connection, ¶ms) - .await? - .into_iter() - .map(|row| row.0) - .collect(); + let rows = crate::span_batches::read_spans(client, connection, params, u64::MAX).await?; if rows.is_empty() { return Ok(None); } @@ -234,6 +230,125 @@ pub async fn get_trace( Ok(resolve_trace(trace_id, &trace_ref, &rows, &spend_rows)) } +#[derive(Deserialize, Serialize)] +struct SpanPosition { + trace_ref: String, + snapshot_ms: u64, + offset: usize, + version: String, +} + +pub async fn get_trace_page( + client: &Client, + connection: &Connection, + access: &ReadAccessParams, + trace_id: &str, + trace_ref: &str, + cursor: Option<&str>, + page_size: u32, +) -> Result, Error> { + if !(1..=500).contains(&page_size) { + return Err(Error::InvalidParameters); + } + let Some(trace_ref) = reference(client, connection, access, trace_id, trace_ref).await? else { + return Ok(None); + }; + let position = match cursor { + Some(cursor) => { + let position: SpanPosition = decode_cursor(cursor, "span")?; + if position.trace_ref != trace_ref || position.snapshot_ms == 0 { + return Err(Error::InvalidCursor("span")); + } + position + } + None => SpanPosition { + trace_ref: trace_ref.clone(), + snapshot_ms: (time::OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) + as u64, + offset: 0, + version: String::new(), + }, + }; + let key_bytes = serde_json::to_vec(&( + connection.url().as_str(), + access, + trace_id, + &trace_ref, + position.snapshot_ms, + )) + .map_err(|_| Error::InvalidParameters)?; + let key = format!("{:x}", Sha256::digest(key_bytes)); + let snapshot = if let Some(trace) = TRACE_SNAPSHOTS.get(&key).await { + trace + } else { + let params = TraceSpansParams { + access: access.clone(), + trace_id: trace_id.to_owned(), + trace_ref: trace_ref.clone(), + }; + let rows = + crate::span_batches::read_spans(client, connection, params, position.snapshot_ms) + .await?; + let spend_rows = spend(client, connection, access, &rows).await; + let Some(trace) = resolve_trace(trace_id, &trace_ref, &rows, &spend_rows) else { + return Ok(None); + }; + let trace = Arc::new(trace); + TRACE_SNAPSHOTS.insert(key, Arc::clone(&trace)).await; + trace + }; + let span_ids: Vec<&str> = snapshot + .spans + .iter() + .map(|span| span.span_id.as_str()) + .collect(); + let version = format!( + "{:x}", + Sha256::digest(serde_json::to_vec(&span_ids).map_err(|_| Error::InvalidResponse)?) + ); + if cursor.is_some() && position.version != version { + return Err(Error::TraceChanged); + } + let mut trace = Trace { + summary: snapshot.summary.clone(), + agents: snapshot.agents.clone(), + spans: Vec::new(), + next_cursor: None, + }; + if position.offset > snapshot.spans.len() { + return Err(Error::InvalidCursor("span")); + } + let end = position + .offset + .saturating_add(page_size as usize) + .min(snapshot.spans.len()); + trace.next_cursor = (end < snapshot.spans.len()).then(|| { + encode_cursor(&SpanPosition { + offset: end, + version: version.clone(), + ..position + }) + }); + trace.spans = snapshot.spans[position.offset..end].to_vec(); + while serde_json::to_vec(&trace) + .map_err(|_| Error::InvalidResponse)? + .len() + > litellm_storage_clickhouse::READ_LIMITS.response_bytes + { + if trace.spans.len() <= 1 { + return Err(Error::ReadTooLarge); + } + trace.spans.truncate(trace.spans.len() / 2); + trace.next_cursor = Some(encode_cursor(&SpanPosition { + trace_ref: trace_ref.clone(), + snapshot_ms: position.snapshot_ms, + offset: position.offset + trace.spans.len(), + version: version.clone(), + })); + } + Ok(Some(trace)) +} + pub async fn get_span( client: &Client, connection: &Connection, diff --git a/litellm-rust/crates/traces-clickhouse/src/span_batches.rs b/litellm-rust/crates/traces-clickhouse/src/span_batches.rs new file mode 100644 index 00000000000..1d903ff0f54 --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/src/span_batches.rs @@ -0,0 +1,181 @@ +use litellm_http::Client; +use litellm_storage_clickhouse::{Query, fetch}; +use litellm_traces::query::named as contracts; +use serde::Serialize; + +use crate::{Connection, Error, query::named::TraceSpansRow}; + +const PAGE_SIZE: u32 = 256; +const MAX_GRAPH_BYTES: usize = 64 * 1024 * 1024; +const MAX_GRAPH_SPANS: usize = 100_000; + +#[derive(Default)] +struct ReadBudget { + bytes: usize, + rows: usize, +} + +impl ReadBudget { + fn reserve(&mut self, bytes: usize) -> Result<(), Error> { + self.bytes = self.bytes.saturating_add(bytes); + if self.bytes > MAX_GRAPH_BYTES || self.rows == MAX_GRAPH_SPANS { + return Err(Error::ReadTooLarge); + } + self.rows += 1; + Ok(()) + } + + fn record(&mut self, row: &impl Serialize) -> Result<(), Error> { + let bytes = serde_json::to_vec(row).map_err(|_| Error::InvalidResponse)?; + self.reserve(bytes.len()) + } +} + +#[derive(Serialize)] +struct Parameters { + #[serde(flatten)] + trace: contracts::TraceSpansParams, + after_span_id: String, + page_size: u32, + snapshot_ms: u64, +} + +struct SpanBatch; + +impl Query for SpanBatch { + type Params = Parameters; + type Row = TraceSpansRow; + + const SQL: &'static str = include_str!("../query/trace_span_batch.sql"); +} + +pub(crate) async fn read_spans( + client: &Client, + connection: &Connection, + trace: contracts::TraceSpansParams, + snapshot_ms: u64, +) -> Result, Error> { + let mut parameters = Parameters { + trace, + after_span_id: String::new(), + page_size: PAGE_SIZE, + snapshot_ms, + }; + let mut spans = Vec::new(); + let mut budget = ReadBudget::default(); + loop { + let page = match fetch::(client, connection, ¶meters).await { + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) + if parameters.page_size > 1 => + { + parameters.page_size /= 2; + continue; + } + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { + return Err(Error::ReadTooLarge); + } + result => result?, + }; + let complete = page.len() < parameters.page_size as usize; + if let Some(last) = page.last() { + parameters.after_span_id.clone_from(&last.0.span_id); + } + for row in page { + budget.record(&row)?; + spans.push(row.0); + } + if complete { + spans.sort_by_key(|row| row.start_ns); + return Ok(spans); + } + parameters.page_size = (parameters.page_size * 2).min(PAGE_SIZE); + } +} + +#[derive(Serialize)] +struct SpendParameters { + #[serde(flatten)] + lookup: crate::query::named::SpendByResponseIdsParams, + has_cursor: u8, + after_team: String, + after_ms: i64, + after_id: String, + page_size: u32, +} + +struct SpendBatch; + +impl Query for SpendBatch { + type Params = SpendParameters; + type Row = crate::query::named::SpendByResponseIdsRow; + + const SQL: &'static str = include_str!("../query/spend_batch.sql"); +} + +pub(crate) async fn read_spend( + client: &Client, + connection: &Connection, + lookup: crate::query::named::SpendByResponseIdsParams, +) -> Result, Error> { + let mut parameters = SpendParameters { + lookup, + has_cursor: 0, + after_team: String::new(), + after_ms: 0, + after_id: String::new(), + page_size: PAGE_SIZE, + }; + let mut rows = Vec::new(); + let mut budget = ReadBudget::default(); + loop { + let page = match fetch::(client, connection, ¶meters).await { + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) + if parameters.page_size > 1 => + { + parameters.page_size /= 2; + continue; + } + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { + return Err(Error::ReadTooLarge); + } + result => result?, + }; + let complete = page.len() < parameters.page_size as usize; + if let Some(last) = page.last() { + parameters.has_cursor = 1; + parameters.after_team.clone_from(&last.0.team_id); + parameters.after_ms = last.0.start_ms; + parameters.after_id.clone_from(&last.0.request_id); + } + for row in page { + budget.record(&row)?; + rows.push(row.0); + } + if complete { + return Ok(rows); + } + parameters.page_size = (parameters.page_size * 2).min(PAGE_SIZE); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use rstest::rstest; + + #[rstest] + #[case::byte_boundary(MAX_GRAPH_BYTES - 1, 0, 1, false)] + #[case::byte_overflow(MAX_GRAPH_BYTES - 1, 0, 2, true)] + #[case::integer_overflow(MAX_GRAPH_BYTES, 0, usize::MAX, true)] + #[case::row_boundary(0, MAX_GRAPH_SPANS - 1, 1, false)] + #[case::row_overflow(0, MAX_GRAPH_SPANS, 1, true)] + fn accumulation_stops_at_the_graph_budget( + #[case] bytes: usize, + #[case] rows: usize, + #[case] next: usize, + #[case] rejected: bool, + ) { + let mut budget = ReadBudget { bytes, rows }; + assert_eq!(budget.reserve(next).is_err(), rejected); + } +} diff --git a/litellm-rust/crates/traces-clickhouse/tests/reads.rs b/litellm-rust/crates/traces-clickhouse/tests/reads.rs new file mode 100644 index 00000000000..ff4461e2157 --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/tests/reads.rs @@ -0,0 +1,436 @@ +use std::collections::BTreeMap; + +use litellm_traces::query::named::ReadAccessParams; +use litellm_traces_clickhouse::{ + Connection, InsertTable, QueryScope, get_trace, get_trace_page, insert_rows, list_traces, +}; +use rstest::rstest; +use serde_json::json; + +#[path = "queries/support.rs"] +mod fixtures; +mod support; + +use fixtures::{DATABASE, SeededDatabase, migrated_database, seeded_database}; +use support::TestResult; + +#[rstest] +#[case::many_runs(50, 21, 0, false)] +#[case::one_large_run(1, 1100, 0, false)] +#[case::large_rows(1, 280, 20_000, false)] +#[case::many_costs(1, 1101, 0, true)] +#[tokio::test] +async fn large_runs_remain_complete_under_default_reader_limits( + #[future(awt)] migrated_database: TestResult, + #[case] runs: usize, + #[case] steps: usize, + #[case] name_bytes: usize, + #[case] costed: bool, +) -> TestResult { + let fixture = migrated_database?; + let client = &fixture.database.client; + let writer = Connection::writer(&fixture.database.url)?; + for run in 0..runs { + let rows = (0..steps) + .map(|step| { + BTreeMap::from([ + ( + "Timestamp".into(), + json!(1_790_000_000_000_000_000_i64 + step as i64), + ), + ("TraceId".into(), json!(format!("trace-{run:04}"))), + ("SpanId".into(), json!(format!("span-{step:04}"))), + ( + "ParentSpanId".into(), + json!(if step == 0 { "" } else { "span-0000" }), + ), + ( + "SpanName".into(), + json!(if name_bytes == 0 { + format!("step-{step}") + } else { + "x".repeat(name_bytes) + }), + ), + ( + "ObservationType".into(), + json!(if step == 0 { + "agent" + } else if costed { + "llm" + } else { + "tool" + }), + ), + ("TeamId".into(), json!("team-a")), + ("ApiKeyHash".into(), json!("key-a")), + ("Duration".into(), json!(1000)), + ( + "LiteLLMRequestId".into(), + json!(if costed && step > 0 { + format!("response-{step}") + } else { + String::new() + }), + ), + ]) + }) + .collect::>(); + for chunk in rows.chunks(100) { + insert_rows( + client, + &writer, + DATABASE, + InsertTable::OtelTraces, + chunk.to_vec(), + ) + .await?; + } + } + if costed { + let costs = (1..steps) + .map(|step| { + BTreeMap::from([ + ("request_id".into(), json!(format!("request-{step}"))), + ("response_id".into(), json!(format!("response-{step}"))), + ("team_id".into(), json!("team-a")), + ("api_key".into(), json!("key-a")), + ("start_time".into(), json!(1_790_000_000_000_i64)), + ("end_time".into(), json!(1_790_000_000_001_i64)), + ("spend".into(), json!(0.25)), + ]) + }) + .collect::>(); + insert_rows(client, &writer, DATABASE, InsertTable::SpendLogs, costs).await?; + } + let reader = fixture + .readers + .connection(client, &QueryScope::All, "fixture-secret") + .await?; + let access = ReadAccessParams { + all_teams: false, + user_id: String::new(), + team_ids: vec!["team-a".into()], + }; + let page = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?; + assert_eq!(page.data.len(), runs); + for summary in &page.data { + assert_eq!(summary.span_count, steps as u64); + assert_eq!( + if costed { + summary.llm_calls + } else { + summary.tool_calls + }, + (steps - 1) as u64 + ); + if costed { + assert_eq!(summary.spend, Some((steps - 1) as f64 * 0.25)); + } + } + let trace_ref = &page + .data + .iter() + .find(|run| run.trace_id == "trace-0000") + .ok_or("missing run")? + .trace_ref; + let detail = get_trace(client, &reader, &access, "trace-0000", trace_ref) + .await? + .ok_or("missing trace")?; + assert_eq!(detail.spans.len(), steps); + assert_eq!(detail.spans[0].span_id, "span-0000"); + assert_eq!( + detail.spans[steps - 1].span_id, + format!("span-{:04}", steps - 1) + ); + assert_eq!( + if costed { + detail.summary.llm_calls + } else { + detail.summary.tool_calls + }, + (steps - 1) as u64 + ); + let mut cursor = None; + let mut ids = Vec::new(); + loop { + let page = get_trace_page( + client, + &reader, + &access, + "trace-0000", + trace_ref, + cursor.as_deref(), + 200, + ) + .await? + .ok_or("missing page")?; + assert_eq!(page.summary, detail.summary); + assert!(page.spans.len() <= 200); + assert!( + serde_json::to_vec(&page)?.len() + <= litellm_storage_clickhouse::READ_LIMITS.response_bytes + ); + ids.extend(page.spans.into_iter().map(|span| span.span_id)); + cursor = page.next_cursor; + if cursor.is_none() { + break; + } + } + assert_eq!( + ids, + detail + .spans + .iter() + .map(|span| span.span_id.clone()) + .collect::>() + ); + let denied = ReadAccessParams { + team_ids: vec!["other-team".into()], + ..access + }; + assert!( + get_trace(client, &reader, &denied, "trace-0000", trace_ref) + .await? + .is_none() + ); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn cursor_pages_keep_a_tenant_scoped_snapshot_when_more_spans_arrive( + #[future(awt)] seeded_database: TestResult, +) -> TestResult { + let fixture = seeded_database?; + let client = &fixture.database.client; + let reader = fixture + .readers + .connection(client, &QueryScope::All, "fixture-secret") + .await?; + let access = ReadAccessParams { + all_teams: true, + user_id: String::new(), + team_ids: Vec::new(), + }; + let listed = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 10).await?; + let summary = listed + .data + .iter() + .find(|summary| summary.span_count == 3) + .ok_or("missing fixture")?; + let first = get_trace_page( + client, + &reader, + &access, + &summary.trace_id, + &summary.trace_ref, + None, + 1, + ) + .await? + .ok_or("missing first page")?; + let original_ids = get_trace( + client, + &reader, + &access, + &summary.trace_id, + &summary.trace_ref, + ) + .await? + .ok_or("missing trace")? + .spans + .into_iter() + .map(|span| span.span_id) + .collect::>(); + let writer = Connection::writer(&fixture.database.url)?; + insert_rows( + client, + &writer, + DATABASE, + InsertTable::OtelTraces, + vec![BTreeMap::from([ + ("Timestamp".into(), json!(1_790_000_000_000_000_000_i64)), + ("TraceId".into(), json!(summary.trace_id)), + ("SpanId".into(), json!("late-span")), + ("ParentSpanId".into(), json!(first.spans[0].span_id)), + ("TeamId".into(), json!("team-a")), + ("ApiKeyHash".into(), json!("key-a")), + ("EngineReceivedMs".into(), json!(u64::MAX / 2)), + ])], + ) + .await?; + let denied = ReadAccessParams { + all_teams: false, + user_id: String::new(), + team_ids: vec!["not-this-team".into()], + }; + assert!( + get_trace_page( + client, + &reader, + &denied, + &summary.trace_id, + &summary.trace_ref, + first.next_cursor.as_deref(), + 1 + ) + .await? + .is_none() + ); + let first_cursor = first.next_cursor.clone(); + let mut cursor = first.next_cursor; + let mut ids = first + .spans + .into_iter() + .map(|span| span.span_id) + .collect::>(); + while let Some(current) = cursor { + let next = get_trace_page( + client, + &reader, + &access, + &summary.trace_id, + &summary.trace_ref, + Some(¤t), + 1, + ) + .await? + .ok_or("missing next page")?; + assert_eq!(next.summary.span_count, 3); + ids.extend(next.spans.into_iter().map(|span| span.span_id)); + cursor = next.next_cursor; + } + assert_eq!(ids, original_ids); + let refreshed = get_trace( + client, + &reader, + &access, + &summary.trace_id, + &summary.trace_ref, + ) + .await? + .ok_or("missing refreshed trace")?; + assert_eq!(refreshed.spans.len(), 4); + assert!(matches!( + get_trace_page( + client, + &reader, + &access, + &summary.trace_id, + &summary.trace_ref, + Some("invalid"), + 1 + ) + .await, + Err(litellm_traces_clickhouse::Error::InvalidCursor("span")) + )); + let backdated = json!({ + "Timestamp": "2026-09-01 00:00:00.000000000", + "TraceId": summary.trace_id, + "SpanId": "backdated-span", + "EngineReceivedMs": 1, + "TeamId": "team-a", + "ApiKeyHash": "key-a" + }); + client + .post(writer.url().clone()) + .body(format!( + "INSERT INTO {DATABASE}.otel_traces FORMAT JSONEachRow\n{backdated}" + )) + .send() + .await? + .error_for_status()?; + let uncached_reader = + Connection::reader(&format!("{}?max_threads=1", fixture.database.url), DATABASE)?; + let changed = get_trace_page( + client, + &uncached_reader, + &access, + &summary.trace_id, + &summary.trace_ref, + first_cursor.as_deref(), + 1, + ) + .await; + assert!( + matches!(changed, Err(litellm_traces_clickhouse::Error::TraceChanged)), + "{changed:?}" + ); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals( + #[future(awt)] seeded_database: TestResult, +) -> TestResult { + let fixture = seeded_database?; + let client = &fixture.database.client; + let reader = fixture + .readers + .connection(client, &QueryScope::All, "fixture-secret") + .await?; + let access = ReadAccessParams { + all_teams: true, + user_id: String::new(), + team_ids: Vec::new(), + }; + let before = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?; + let run = before + .data + .iter() + .find(|run| run.span_count == 3) + .ok_or("missing fixture")?; + let writer = Connection::writer(&fixture.database.url)?; + insert_rows( + client, + &writer, + DATABASE, + InsertTable::OtelTraces, + vec![BTreeMap::from([ + ("Timestamp".into(), json!(1_790_000_000_000_000_000_i64)), + ("TraceId".into(), json!(run.trace_id)), + ("SpanId".into(), json!("oversized-child")), + ("ParentSpanId".into(), json!("0101010101010101")), + ( + "SpanName".into(), + json!("x".repeat(litellm_storage_clickhouse::READ_LIMITS.response_bytes + 1)), + ), + ("ObservationType".into(), json!("tool")), + ("TeamId".into(), json!("team-a")), + ("ApiKeyHash".into(), json!("key-a")), + ])], + ) + .await?; + let after = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?; + assert_eq!(after.data.len(), before.data.len()); + let limited = after + .data + .iter() + .find(|item| item.trace_ref == run.trace_ref) + .ok_or("missing run")?; + assert!(limited.resolution_limited); + assert_eq!(limited.span_count, 4); + assert!( + after + .data + .iter() + .filter(|item| item.trace_ref != run.trace_ref) + .all(|item| !item.resolution_limited) + ); + assert!(matches!( + get_trace_page( + client, + &reader, + &access, + &run.trace_id, + &run.trace_ref, + None, + 200 + ) + .await, + Err(litellm_traces_clickhouse::Error::ReadTooLarge) + )); + Ok(()) +} diff --git a/litellm-rust/crates/traces/src/resolve/view.rs b/litellm-rust/crates/traces/src/resolve/view.rs index 51145e6bbbf..9c1c3756856 100644 --- a/litellm-rust/crates/traces/src/resolve/view.rs +++ b/litellm-rust/crates/traces/src/resolve/view.rs @@ -171,6 +171,7 @@ pub fn resolve_trace( .map(|(_, (span, _))| span.input_preview.clone()) .unwrap_or_default(); let summary = TraceSummary { + resolution_limited: false, trace_id: trace_id.to_owned(), trace_ref: trace_ref.to_owned(), name: spans[root].name.clone(), @@ -209,11 +210,13 @@ pub fn resolve_trace( summary, agents, spans, + next_cursor: None, }) } pub fn listed_summary(row: &ListTracesRow) -> TraceSummary { TraceSummary { + resolution_limited: true, trace_id: row.trace_id.clone(), trace_ref: row.trace_ref.clone(), name: row.name.clone(), diff --git a/litellm-rust/crates/traces/src/view.rs b/litellm-rust/crates/traces/src/view.rs index b7a67ac6822..a864c740526 100644 --- a/litellm-rust/crates/traces/src/view.rs +++ b/litellm-rust/crates/traces/src/view.rs @@ -17,7 +17,7 @@ pub enum SpanStatus { } #[macro_rules_attribute::apply(response_type)] -#[derive(Debug, PartialEq)] +#[derive(Clone, Debug, PartialEq)] pub struct Span { pub span_id: String, pub parent_span_id: Option, @@ -41,7 +41,7 @@ pub struct Span { /// One distinct agent in a trace: 200 invocations of `researcher` are one node. #[macro_rules_attribute::apply(response_type)] -#[derive(Debug, PartialEq)] +#[derive(Clone, Debug, PartialEq)] pub struct AgentNode { pub name: String, pub parent_agent: Option, @@ -53,8 +53,10 @@ pub struct AgentNode { } #[macro_rules_attribute::apply(response_type)] -#[derive(Debug, PartialEq)] +#[derive(Clone, Debug, PartialEq)] pub struct TraceSummary { + #[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))] + pub resolution_limited: bool, pub trace_id: String, #[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))] pub trace_ref: String, @@ -81,11 +83,13 @@ pub struct TraceSummary { } #[macro_rules_attribute::apply(response_type)] -#[derive(Debug, PartialEq)] +#[derive(Clone, Debug, PartialEq)] pub struct Trace { pub summary: TraceSummary, pub agents: Vec, pub spans: Vec, + #[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))] + pub next_cursor: Option, } #[macro_rules_attribute::apply(response_type)] diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index 2cc8e035ebd..c27df51ade2 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -290,6 +290,7 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset( "LiteLLM_Config", "LiteLLM_SpendLogs", "LiteLLM_BudgetWindowSpend", + "LiteLLM_BackgroundInteractionSettlement", "LiteLLM_ErrorLogs", "LiteLLM_UserNotifications", "LiteLLM_TeamMembership", diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index 50c6e80b234..1563e4b5b55 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -140,7 +140,7 @@ async def list_agent_traces( context: Annotated[TraceAccessContext, Depends(provide_trace_access)], start_ms: Annotated[int | None, Query(description="Window start, unix ms. Default: 24h ago")] = None, end_ms: Annotated[int | None, Query(description="Window end, unix ms. Default: now")] = None, - cursor: Annotated[str | None, Query()] = None, + cursor: Annotated[str | None, Query(max_length=512)] = None, ) -> TracePage: now_ms: Final = int(time.time() * 1000) try: @@ -153,6 +153,13 @@ async def list_agent_traces( ) except ValueError as error: raise HTTPException(status_code=400, detail=str(error)) from error + except OverflowError as error: + raise HTTPException( + status_code=413, detail="Trace is too large for this view. Use a filtered trace query." + ) from error + except RuntimeError as error: + verbose_proxy_logger.warning("Trace read unavailable: %s", error) + raise HTTPException(status_code=503, detail="Traces are temporarily unavailable. Please try again.") from error class TraceQueryRequest(BaseModel): @@ -228,12 +235,21 @@ async def get_agent_trace( trace_id: str, context: Annotated[TraceAccessContext, Depends(provide_trace_access)], trace_ref: Annotated[str, Query()] = "", + cursor: Annotated[str | None, Query(max_length=512)] = None, + page_size: Annotated[int | None, Query(ge=1, le=500)] = None, ) -> Trace: tracing, scope = context.reader() try: - trace: Final = await tracing.get_trace(trace_id, scope, trace_ref) + trace: Final = await tracing.get_trace(trace_id, scope, trace_ref, cursor, page_size) except ValueError as error: raise HTTPException(status_code=400, detail=str(error)) from error + except OverflowError as error: + raise HTTPException( + status_code=413, detail="Trace is too large for this view. Use a filtered trace query." + ) from error + except RuntimeError as error: + verbose_proxy_logger.warning("Trace read unavailable: %s", error) + raise HTTPException(status_code=503, detail="Traces are temporarily unavailable. Please try again.") from error if trace is None: raise HTTPException(status_code=404, detail=f"Trace {trace_id} not found") return trace @@ -251,6 +267,13 @@ async def get_agent_trace_span( span: Final = await tracing.get_span(trace_id, span_id, scope, trace_ref) except ValueError as error: raise HTTPException(status_code=400, detail=str(error)) from error + except OverflowError as error: + raise HTTPException( + status_code=413, detail="Trace is too large for this view. Use a filtered trace query." + ) from error + except RuntimeError as error: + verbose_proxy_logger.warning("Trace read unavailable: %s", error) + raise HTTPException(status_code=503, detail="Traces are temporarily unavailable. Please try again.") from error if span is None: raise HTTPException(status_code=404, detail=f"Span {span_id} not found") return span @@ -269,6 +292,13 @@ async def get_agent_trace_span_error( page: Final = await tracing.get_span_error(trace_id, span_id, scope, trace_ref, cursor) except ValueError as error: raise HTTPException(status_code=400, detail=str(error)) from error + except OverflowError as error: + raise HTTPException( + status_code=413, detail="Trace is too large for this view. Use a filtered trace query." + ) from error + except RuntimeError as error: + verbose_proxy_logger.warning("Trace read unavailable: %s", error) + raise HTTPException(status_code=503, detail="Traces are temporarily unavailable. Please try again.") from error if page is None: raise HTTPException(status_code=404, detail="Span diagnostic not found or no longer available") return page diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index e7ecec4df0f..c146a6eac92 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -45,7 +45,9 @@ class NativeTraceStorage: def list_traces( self, scope: TraceScope, start_ms: int, end_ms: int, cursor: str | None, limit: int ) -> Future[JsonValue]: ... - def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str) -> Future[JsonValue]: ... + def get_trace( + self, trace_id: str, scope: TraceScope, trace_ref: str, cursor: str | None = None, page_size: int | None = None + ) -> Future[JsonValue]: ... def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str) -> Future[JsonValue]: ... def get_span_error( self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str, cursor: str | None diff --git a/litellm/rust_bridge/trace/generated/types.py b/litellm/rust_bridge/trace/generated/types.py index e1c09ffe281..127e86e9160 100644 --- a/litellm/rust_bridge/trace/generated/types.py +++ b/litellm/rust_bridge/trace/generated/types.py @@ -99,6 +99,7 @@ class UIMessage(typing_extensions.TypedDict): class TraceSummary(typing_extensions.TypedDict): + resolution_limited: ReadOnly[NotRequired[bool]] trace_id: ReadOnly[str] trace_ref: ReadOnly[NotRequired[str]] name: ReadOnly[str] @@ -145,6 +146,7 @@ class Trace(typing_extensions.TypedDict): summary: ReadOnly[TraceSummary] agents: ReadOnly[tuple[AgentNode, ...]] spans: ReadOnly[tuple[Span, ...]] + next_cursor: ReadOnly[NotRequired[str | None]] class TracePage(typing_extensions.TypedDict): diff --git a/litellm/rust_bridge/trace/storage.py b/litellm/rust_bridge/trace/storage.py index 7e7c519365d..a80e41a5aa3 100644 --- a/litellm/rust_bridge/trace/storage.py +++ b/litellm/rust_bridge/trace/storage.py @@ -67,7 +67,9 @@ class NativeStore(Protocol): self, scope: TraceScope, start_ms: int, end_ms: int, cursor: str | None, limit: int ) -> Awaitable[JsonValue]: ... - def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str) -> Awaitable[JsonValue]: ... + def get_trace( + self, trace_id: str, scope: TraceScope, trace_ref: str, cursor: str | None = None, page_size: int | None = None + ) -> Awaitable[JsonValue]: ... def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str) -> Awaitable[JsonValue]: ... @@ -189,8 +191,15 @@ class ClickHouseStorage: result: Final = await self._native.list_traces(scope, start_ms, end_ms, cursor, limit) return _validate_query_response(_TRACE_PAGE, result) - async def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str = "") -> Trace | None: - result: Final = await self._native.get_trace(trace_id, scope, trace_ref) + async def get_trace( + self, + trace_id: str, + scope: TraceScope, + trace_ref: str = "", + cursor: str | None = None, + page_size: int | None = None, + ) -> Trace | None: + result: Final = await self._native.get_trace(trace_id, scope, trace_ref, cursor, page_size) return _validate_query_response(_TRACE, result) async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None: diff --git a/litellm/tracing/receiver.py b/litellm/tracing/receiver.py index 8df2af4f064..cf6bf8feb01 100644 --- a/litellm/tracing/receiver.py +++ b/litellm/tracing/receiver.py @@ -99,8 +99,15 @@ class TraceReceiver: async def list_traces(self, scope: TraceScope, start_ms: int, end_ms: int, cursor: str | None = None) -> TracePage: return await self.storage.list_traces(scope, start_ms, end_ms, cursor, AGENT_TRACING_LIST_PAGE_SIZE) - async def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str = "") -> Trace | None: - return await self.storage.get_trace(trace_id, scope, trace_ref) + async def get_trace( + self, + trace_id: str, + scope: TraceScope, + trace_ref: str = "", + cursor: str | None = None, + page_size: int | None = None, + ) -> Trace | None: + return await self.storage.get_trace(trace_id, scope, trace_ref, cursor, page_size) async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None: return await self.storage.get_span(trace_id, span_id, scope, trace_ref) diff --git a/scripts/trace_codegen/schemas/traces/Trace.json b/scripts/trace_codegen/schemas/traces/Trace.json index ee02937f556..dc66f77412c 100644 --- a/scripts/trace_codegen/schemas/traces/Trace.json +++ b/scripts/trace_codegen/schemas/traces/Trace.json @@ -245,6 +245,10 @@ "minimum": 0, "type": "integer" }, + "resolution_limited": { + "type": "boolean", + "x-python-optional": true + }, "service": { "type": "string" }, @@ -282,6 +286,7 @@ } }, "required": [ + "resolution_limited", "trace_id", "trace_ref", "name", @@ -314,6 +319,13 @@ }, "type": "array" }, + "next_cursor": { + "type": [ + "string", + "null" + ], + "x-python-optional": true + }, "spans": { "items": { "$ref": "#/$defs/Span" @@ -327,7 +339,8 @@ "required": [ "summary", "agents", - "spans" + "spans", + "next_cursor" ], "title": "Trace", "type": "object" diff --git a/scripts/trace_codegen/schemas/traces/TracePage.json b/scripts/trace_codegen/schemas/traces/TracePage.json index 72b2c2b2d95..dfa02ba0d39 100644 --- a/scripts/trace_codegen/schemas/traces/TracePage.json +++ b/scripts/trace_codegen/schemas/traces/TracePage.json @@ -76,6 +76,10 @@ "minimum": 0, "type": "integer" }, + "resolution_limited": { + "type": "boolean", + "x-python-optional": true + }, "service": { "type": "string" }, @@ -113,6 +117,7 @@ } }, "required": [ + "resolution_limited", "trace_id", "trace_ref", "name", diff --git a/tests/test_litellm/tracing/test_receiver.py b/tests/test_litellm/tracing/test_receiver.py index 66b971e8e63..70282fcf694 100644 --- a/tests/test_litellm/tracing/test_receiver.py +++ b/tests/test_litellm/tracing/test_receiver.py @@ -61,11 +61,12 @@ async def test_ingest_rejects_oversized_body_before_storage() -> None: @pytest.mark.asyncio -async def test_reads_delegate_to_storage() -> None: +@pytest.mark.parametrize("cursor,page_size", ((None, None), ("next", 200))) +async def test_reads_delegate_to_storage(cursor: str | None, page_size: int | None) -> None: storage: Final = _fake_storage() scope: Final[TraceScope] = {"all_teams": 0, "user_id": "", "team_ids": ("team-research",)} - assert await TraceReceiver(storage).get_trace("t1", scope) is None - storage.get_trace.assert_awaited_once_with("t1", scope, "") + assert await TraceReceiver(storage).get_trace("t1", scope, "", cursor, page_size) is None + storage.get_trace.assert_awaited_once_with("t1", scope, "", cursor, page_size) @pytest.mark.asyncio diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index 16d02865e92..cff67d1f57d 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -266,7 +266,7 @@ def test_get_trace_404_and_200(client, receiver): response = client.get("/v1/traces/t1") assert response.status_code == 200 assert response.json() == TRACE_RESPONSE - receiver.get_trace.assert_awaited_with("t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "") + receiver.get_trace.assert_awaited_with("t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "", None, None) def test_get_span_404_and_200(client, receiver): @@ -278,10 +278,45 @@ def test_get_span_404_and_200(client, receiver): receiver.get_span.assert_awaited_with("t1", "s1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "") -def test_trace_detail_passes_scoped_reference(client, receiver): +@pytest.mark.parametrize("suffix,cursor,page_size", [("", None, None), ("&cursor=next&page_size=200", "next", 200)]) +def test_trace_detail_passes_scoped_reference(client, receiver, suffix, cursor, page_size): receiver.get_trace.return_value = TRACE_RESPONSE - assert client.get("/v1/traces/t1?trace_ref=run-one").status_code == 200 - receiver.get_trace.assert_awaited_with("t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one") + assert client.get(f"/v1/traces/t1?trace_ref=run-one{suffix}").status_code == 200 + receiver.get_trace.assert_awaited_with( + "t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one", cursor, page_size + ) + + +@pytest.mark.parametrize( + "path,method", + ( + ("/v1/traces", "list_traces"), + ("/v1/traces/t1", "get_trace"), + ("/v1/traces/t1/spans/s1", "get_span"), + ("/v1/traces/t1/spans/s1/error", "get_span_error"), + ), +) +@pytest.mark.parametrize( + "error,status,message", + ( + (RuntimeError("private database details"), 503, "Traces are temporarily unavailable. Please try again."), + (OverflowError("private query details"), 413, "Trace is too large for this view. Use a filtered trace query."), + ), +) +def test_read_failures_are_actionable_without_exposing_database_details( + client: TestClient, receiver: MagicMock, path: str, method: str, error: Exception, status: int, message: str +) -> None: + getattr(receiver, method).side_effect = error + response: Final = client.get(path) + assert response.status_code == status + assert response.json() == {"detail": message} + + +@pytest.mark.parametrize("query", ("page_size=0", "page_size=501", "cursor=" + "x" * 513)) +def test_trace_page_rejects_unbounded_parameters(client: TestClient, receiver: MagicMock, query: str) -> None: + response: Final = client.get(f"/v1/traces/t1?{query}") + assert response.status_code == 422 + receiver.get_trace.assert_not_awaited() def test_invalid_export_and_cursor_are_client_errors(client, receiver): @@ -370,7 +405,9 @@ def test_injected_receiver_ingests_with_the_authenticated_tenant(client: TestCli storage: Final = MagicMock(spec=ClickHouseStorage) storage.ingest = AsyncMock(return_value=1) client.app.dependency_overrides[tracing_endpoints.provide_receiver] = lambda: TraceReceiver(storage) - response: Final = client.post("/v1/traces", content=b'{"resourceSpans": []}', headers={"content-type": "application/json"}) + response: Final = client.post( + "/v1/traces", content=b'{"resourceSpans": []}', headers={"content-type": "application/json"} + ) assert response.status_code == 200, response.text assert response.json() == {} storage.ingest.assert_awaited_once_with( @@ -538,9 +575,7 @@ def test_sql_and_help_use_authenticated_scope( result: Final = client.post("/v1/traces/query", json={"sql": "SELECT * FROM otel_traces"}) assert result.status_code == 200, result.text assert result.json() == SQL_ENVELOPE - receiver.storage.query_sql.assert_awaited_once_with( - "SELECT * FROM otel_traces", expected_scope, "test-secret" - ) + receiver.storage.query_sql.assert_awaited_once_with("SELECT * FROM otel_traces", expected_scope, "test-secret") help_result: Final = client.get("/v1/traces/query/help") assert help_result.status_code == 200, help_result.text assert help_result.json() == QUERY_HELP @@ -692,7 +727,7 @@ class _NativeConfig: class _NativeReturningHelp(ModuleType): - def __init__(self, help_payload: Mapping[str, object]) -> None: + def __init__(self, help_payload: Mapping[str, object], trace_payload: Mapping[str, object] | None = None) -> None: super().__init__("native_traces") class Storage: @@ -702,12 +737,31 @@ class _NativeReturningHelp(ModuleType): async def query_help(self, scope: AllQueryScope, secret: str) -> Mapping[str, object]: return help_payload + get_trace = AsyncMock(return_value=trace_payload) + + self.trace_read: Final = Storage.get_trace self.NativeTraceConfig: Final = _NativeConfig self.NativeTraceStorage: Final = Storage self.trace_encode_error: Final = bytes self.trace_span_rows: Final = list +@pytest.mark.parametrize("cursor,page_size", ((None, None), ("next", 200))) +async def test_storage_preserves_page_cursor_and_normalizes_native_trace_data( + monkeypatch: pytest.MonkeyPatch, cursor: str | None, page_size: int | None +) -> None: + native: Final = _NativeReturningHelp(QUERY_HELP, {**TRACE_RESPONSE, "next_cursor": "more"}) + monkeypatch.setattr(loader, "_cached_bridge", native) + storage: Final = ClickHouseStorage(TraceStorageConfig("http://clickhouse:8123")) + scope: Final[TraceScope] = {"all_teams": 0, "user_id": "owner", "team_ids": ()} + trace: Final = await storage.get_trace("t1", scope, "run", cursor, page_size) + assert trace is not None + assert trace["next_cursor"] == "more" + assert trace["spans"] == () + assert trace["summary"]["span_count"] == 0 + native.trace_read.assert_awaited_once_with("t1", scope, "run", cursor, page_size) + + async def test_storage_validates_the_native_query_help_value(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(loader, "_cached_bridge", _NativeReturningHelp(QUERY_HELP)) storage: Final = ClickHouseStorage(TraceStorageConfig("http://clickhouse:8123")) diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 4cd3b271c66..52021f61ba6 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -1975,10 +1975,15 @@ export const agentTraceListCall = async ({ export const sendOtlpTraceCall = async (accessToken: string, exportRequest: object): Promise => apiClient.post(`/v1/traces`, { accessToken, body: exportRequest }); -export const agentTraceCall = async (accessToken: string, traceId: string, traceRef?: string): Promise => +export const agentTraceCall = async ( + accessToken: string, + traceId: string, + traceRef?: string, + cursor?: string | null, +): Promise => apiClient.get(`/v1/traces/${encodeURIComponent(traceId)}`, { accessToken, - query: { trace_ref: traceRef || undefined }, + query: { trace_ref: traceRef || undefined, cursor: cursor ?? undefined, page_size: 200 }, }); export const agentTraceSpanCall = async ( diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.integration.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.integration.test.tsx index 816bf5f9267..b613d028f54 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.integration.test.tsx @@ -85,6 +85,25 @@ describe("AgentTracesSection", () => { vi.mocked(apiClient.get).mockResolvedValue({ data: [] }); }); + it("keeps the original time window and loaded rows when another page fails", async () => { + const user = userEvent.setup(); + const now = vi.spyOn(Date, "now").mockReturnValue(Date.parse("2026-10-01T00:00Z")); + vi.mocked(agentTraceListCall).mockResolvedValueOnce({ data: runs.slice(0, 1), next_cursor: "next" }); + renderSection(); + expect(await screen.findByTestId("agent-trace-row")).toBeVisible(); + const first = vi.mocked(agentTraceListCall).mock.calls[0][0]; + now.mockReturnValue(Date.parse("2026-10-01T01:00Z")); + vi.mocked(agentTraceListCall).mockRejectedValue(new ApiError("Please try again", 403, {})); + await user.click(screen.getByRole("button", { name: "Load more" })); + expect(await screen.findByRole("alert")).toHaveTextContent("Could not load more runs"); + expect(screen.getAllByTestId("agent-trace-row")).toHaveLength(1); + expect(vi.mocked(agentTraceListCall).mock.calls[1][0]).toEqual({ ...first, cursor: "next" }); + vi.mocked(agentTraceListCall).mockResolvedValueOnce({ data: runs.slice(1, 2), next_cursor: null }); + await user.click(screen.getByRole("button", { name: "Retry" })); + await waitFor(() => expect(screen.getAllByTestId("agent-trace-row")).toHaveLength(2)); + expect(screen.queryByRole("alert")).not.toBeInTheDocument(); + }); + it.each([ [401, "Your session is no longer valid. Sign out and sign in again."], [403, "Your account does not have access to these traces."], diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx index 45f1b223088..211f50dcd38 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.tsx @@ -229,6 +229,8 @@ export function AgentTracesSection({ isLoading={traces.isLoading || (checkHistory && history.isLoading)} error={traces.error} hasMore={traces.hasMore} + isFetching={traces.isFetching} + onRetry={traces.hasMore ? traces.loadMore : traces.refetch} onLoadMore={traces.loadMore} onOpenTrace={toggleRun} selectedKey={openTrace === null ? null : runKey(openTrace)} diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx index 2158e51dd5b..964b6e82469 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx @@ -16,6 +16,8 @@ interface AgentTracesTableProps { isLoading: boolean; error: Error | null; hasMore: boolean; + isFetching?: boolean; + onRetry?: () => void; onLoadMore: () => void; onOpenTrace: (trace: TraceSummary) => void; selectedKey?: string | null; @@ -53,6 +55,8 @@ export function AgentTracesTable({ isLoading, error, hasMore, + isFetching = false, + onRetry, onLoadMore, onOpenTrace, selectedKey = null, @@ -105,6 +109,14 @@ export function AgentTracesTable({ {firstLine(previewText(run.input_preview)) || traceDisplayName(run)} + {run.resolution_limited && ( + + Partial totals + + )} {run.trace_id} @@ -132,15 +144,24 @@ export function AgentTracesTable({ {isLoading &&
Loading runs…
} {error && ( -
Could not load runs: {error.message}
+
+ + {traces.length ? "Could not load more runs" : "Could not load runs"}: {error.message} + + {onRetry && ( + + )} +
)} {isEmpty && (
No runs match these filters.
)} {hasMore && (
-
)} diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx index 5dccf5cd847..d055426257c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx @@ -155,13 +155,46 @@ describe("RunView", () => { expect(screen.queryByRole("button", { name: "Show details" })).not.toBeInTheDocument(); }); + it("keeps loaded steps and totals after a page fails, then retries the same cursor", async () => { + const user = userEvent.setup(); + const summary = { ...research.summary, span_count: 2 }; + const first: Trace = { ...research, summary, spans: research.spans.slice(0, 1), next_cursor: "next-page" }; + const second: Trace = { + ...research, + summary, + spans: [ + { ...research.spans[1], type: "tool", name: "later-page-tool", parent_span_id: research.spans[0].span_id }, + ], + next_cursor: null, + }; + vi.mocked(agentTraceCall).mockReset(); + vi.mocked(agentTraceCall) + .mockResolvedValueOnce(first) + .mockRejectedValueOnce(new Error("offline")) + .mockResolvedValueOnce(second); + renderWithProviders(); + + expect(await screen.findByText("Showing 1 of 2 steps")).toBeVisible(); + const before = screen.getByRole("banner").textContent; + await user.click(screen.getByRole("button", { name: "Load more steps" })); + expect(await screen.findByText("Could not load more steps. Your loaded steps are still available.")).toBeVisible(); + expect(screen.getByRole("tree", { name: "Spans in time order" })).toHaveTextContent(research.spans[0].name); + await user.click(screen.getByRole("button", { name: "Retry" })); + expect(await screen.findByText("later-page-tool")).toBeVisible(); + expect(screen.getByRole("tree", { name: "Spans in time order" })).toHaveTextContent(research.spans[0].name); + expect(screen.getAllByRole("treeitem")).toHaveLength(2); + expect(screen.getByRole("banner")).toHaveTextContent(before ?? ""); + expect(screen.queryByRole("button", { name: "Load more steps" })).not.toBeInTheDocument(); + expect(vi.mocked(agentTraceCall).mock.calls.map((call) => call[3])).toEqual([null, "next-page", "next-page"]); + }); + it("keeps a way back to the runs table when a run fails to load", async () => { const user = userEvent.setup(); const onBack = vi.fn(); - vi.mocked(agentTraceCall).mockRejectedValue(new Error("trace exceeds the 1000 span read limit")); + vi.mocked(agentTraceCall).mockRejectedValue(new Error("Traces are temporarily unavailable")); renderWithProviders(); - expect(await screen.findByText("trace exceeds the 1000 span read limit")).toBeInTheDocument(); + expect(await screen.findByText("Traces are temporarily unavailable")).toBeInTheDocument(); await user.click(screen.getByRole("button", { name: /back to traces/i })); expect(onBack).toHaveBeenCalledTimes(1); }); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx index acdaef64873..6e8464b797e 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx @@ -2,7 +2,7 @@ import { useLensDemo } from "@/components/lens/LensDemoContext"; import { useTracesApi } from "@/components/lens/services"; -import { useQuery } from "@tanstack/react-query"; +import { useInfiniteQuery } from "@tanstack/react-query"; import { ArrowLeft, Check, Copy } from "lucide-react"; import { useCallback, useEffect, useMemo, useState } from "react"; @@ -368,16 +368,34 @@ interface RunViewProps { embedded?: boolean; } -/** One agent run: header with totals and "Copy for agent", span tree on the left, span details on the right. */ +function initialSpanMissing(trace: Trace | undefined, spanId?: string): boolean { + return Boolean(spanId && trace && !trace.spans.some((span) => span.span_id === spanId)); +} + export function RunView({ traceId, traceRef, initialSpanId, accessToken, onBack, embedded = false }: RunViewProps) { const traces = useTracesApi(accessToken); const [view, setView] = useState("steps"); - const traceQuery = useQuery({ + const traceQueryOptions = { queryKey: ["agentTrace", traceId, traceRef, accessToken], - queryFn: () => traces.trace(traceId, traceRef), + queryFn: ({ pageParam }: { pageParam: string | null }) => traces.trace(traceId, traceRef, pageParam), + initialPageParam: null as string | null, + getNextPageParam: (lastPage: Trace) => lastPage.next_cursor ?? undefined, staleTime: 30_000, - }); - const trace = traceQuery.data; + retry: false, + }; + const traceQuery = useInfiniteQuery(traceQueryOptions); + const trace = useMemo(() => { + const pages = traceQuery.data?.pages; + if (!pages?.length) return undefined; + return { ...pages[0], spans: pages.flatMap((page) => page.spans) }; + }, [traceQuery.data]); + const seekingSpan = initialSpanMissing(trace, initialSpanId); + const { hasNextPage, isFetching, isError, fetchNextPage } = traceQuery; + const canSeek = seekingSpan && hasNextPage; + useEffect(() => { + if (canSeek && !isFetching && !isError) void fetchNextPage(); + }, [canSeek, isFetching, isError, fetchNextPage]); + const pageAction = isError ? "Retry" : "Load more steps"; if (traceQuery.isLoading) { return ( @@ -396,7 +414,7 @@ export function RunView({ traceId, traceRef, initialSpanId, accessToken, onBack, ); } - if (traceQuery.isError || !trace) { + if (!trace) { return (
); } @@ -422,8 +443,35 @@ export function RunView({ traceId, traceRef, initialSpanId, accessToken, onBack, data-testid="run-view" > + {(traceQuery.hasNextPage || traceQuery.isError) && ( +
+ + {traceQuery.isError + ? "Could not load more steps. Your loaded steps are still available." + : `Showing ${trace.spans.length.toLocaleString()} of ${trace.summary.span_count.toLocaleString()} steps`} + + {traceQuery.isError && ( + + )} + +
+ )} ; anyRecorded(): Promise; - trace(traceId: string, traceRef?: string): Promise; + trace(traceId: string, traceRef?: string, cursor?: string | null): Promise; span(traceId: string, spanId: string, traceRef?: string): Promise; spanError( traceId: string, @@ -32,7 +32,7 @@ export function liveTracesApi(accessToken: string): TracesApi { const page = await apiClient.get("/v1/traces", { accessToken, query: { start_ms: 0 } }); return page.data.length > 0; }, - trace: (traceId, traceRef) => agentTraceCall(accessToken, traceId, traceRef), + trace: (traceId, traceRef, cursor) => agentTraceCall(accessToken, traceId, traceRef, cursor), span: (traceId, spanId, traceRef) => agentTraceSpanCall(accessToken, traceId, spanId, traceRef), spanError: (traceId, spanId, options) => agentTraceSpanErrorCall(accessToken, traceId, spanId, options), }; diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/useAgentTraces.ts b/ui/litellm-dashboard/src/components/view_logs/TraceView/useAgentTraces.ts index dd3bbc873b3..b1e20748770 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/useAgentTraces.ts +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/useAgentTraces.ts @@ -7,6 +7,11 @@ import { ApiError } from "@/lib/http/client"; import { LIVE_TAIL_INTERVAL_MS } from "../log_filter_logic"; import type { TracePage, TraceSummary } from "./traceTypes"; +import type { TraceWindow } from "./tracesApi"; + +interface LoadedTracePage extends TracePage { + window: TraceWindow; +} export const TRACING_NOT_ENABLED_STATUS = 501; /** A proxy without the tracing routes at all answers 404; treat it like tracing being off. */ @@ -55,7 +60,7 @@ export const traceWindowStartMs = (startTime: string, endTime: string, isCustomD /** * GET /v1/traces for the Logs page time range, cursor-paginated ("Load more"). - * Preset ranges re-read "now" on every fetch, moving both bounds so the window keeps its length. + * Preset ranges roll on refresh; subsequent pages keep the first page's window. */ export function useAgentTraces({ accessToken, @@ -66,19 +71,20 @@ export function useAgentTraces({ enabled, }: UseAgentTracesOptions): AgentTracesResult { const traces = useTracesApi(accessToken); - const fetchPage = (pageParam: unknown): Promise => { + const fetchPage = async (pageParam: unknown): Promise => { const nowMs = Date.now(); - return traces.list({ + const window = (pageParam as TraceWindow | null) ?? { startMs: traceWindowStartMs(startTime, endTime, isCustomDate, nowMs), endMs: isCustomDate ? moment(endTime).valueOf() : nowMs, - cursor: pageParam as string | null, - }); + }; + return { ...(await traces.list(window)), window }; }; - const queryOptions: Parameters>[0] = { + const queryOptions: Parameters>[0] = { queryKey: ["agentTraces", accessToken, startTime, endTime, isCustomDate], queryFn: ({ pageParam }) => fetchPage(pageParam), initialPageParam: null, - getNextPageParam: (lastPage) => lastPage.next_cursor ?? undefined, + getNextPageParam: (lastPage) => + lastPage.next_cursor ? { ...lastPage.window, cursor: lastPage.next_cursor } : undefined, enabled, staleTime: LIVE_TAIL_INTERVAL_MS, retry: (failureCount, error) => !requiresUserAction(error) && failureCount < 1, @@ -87,7 +93,7 @@ export function useAgentTraces({ refetchOnReconnect: (q) => !requiresUserAction(q.state.error), refetchIntervalInBackground: false, }; - const query = useInfiniteQuery(queryOptions); + const query = useInfiniteQuery(queryOptions); const loaded = useMemo(() => query.data?.pages.flatMap((page) => page.data) ?? [], [query.data]); const notEnabled = isTracingNotEnabled(query.error); @@ -99,7 +105,9 @@ export function useAgentTraces({ notEnabledDetail: notEnabled ? query.error?.message || "Agent tracing is not enabled" : null, error: notEnabled ? null : displayError(query.error), hasMore: query.hasNextPage, - loadMore: () => void query.fetchNextPage(), + loadMore: () => { + if (!query.isFetching) void query.fetchNextPage(); + }, refetch: () => void query.refetch(), }; } diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 2c14143cfe7..3ed4103a227 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -47047,6 +47047,8 @@ export interface components { Trace: { /** Agents */ agents: components["schemas"]["AgentNode"][]; + /** Next Cursor */ + next_cursor?: string | null; /** Spans */ spans: components["schemas"]["Span"][]; summary: components["schemas"]["TraceSummary"]; @@ -47279,6 +47281,8 @@ export interface components { name: string; /** Output Tokens */ output_tokens: number; + /** Resolution Limited */ + resolution_limited?: boolean; /** Service */ service: string; /** Span Count */ @@ -80234,6 +80238,8 @@ export interface operations { parameters: { query?: { trace_ref?: string; + cursor?: string | null; + page_size?: number | null; }; header?: never; path: { From f0abe1bea1feb5978293b21b68fa4849cb13a6a6 Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Sat, 3 Oct 2026 10:56:48 -0700 Subject: [PATCH 05/82] feat(lens): live dot field timeline and full-screen traces view (#44390) * feat(lens): add dot field layout and live status helpers for the traces timeline * test(lens): cover dot field layout, agent colors and live status * feat(lens): draw the traces timeline as a live dot field * feat(lens): add the sweep animation for the live traces timeline * feat(lens): put the lens tabs in a compact header and fill the screen with traces --- ui/litellm-dashboard/src/app/globals.css | 9 + .../src/components/lens/LensWorkspace.tsx | 38 +-- .../view_logs/TraceView/TracesTimeline.tsx | 261 ++++++++++++++---- .../view_logs/TraceView/lensField.test.ts | 114 ++++++++ .../view_logs/TraceView/lensField.ts | 67 +++++ 5 files changed, 423 insertions(+), 66 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/view_logs/TraceView/lensField.test.ts create mode 100644 ui/litellm-dashboard/src/components/view_logs/TraceView/lensField.ts diff --git a/ui/litellm-dashboard/src/app/globals.css b/ui/litellm-dashboard/src/app/globals.css index e1ff921648f..72d304c0dc7 100644 --- a/ui/litellm-dashboard/src/app/globals.css +++ b/ui/litellm-dashboard/src/app/globals.css @@ -473,3 +473,12 @@ [data-slot="dialog-content"][data-nested-dialog-open] { visibility: hidden; } + +@keyframes lens-sweep { + from { + left: -6rem; + } + to { + left: 100%; + } +} diff --git a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx index 0a3d3ae6a0d..a0ac3e3e59c 100644 --- a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx @@ -69,29 +69,31 @@ function LensContent({ const openDemo = onDemo ? () => onDemo(activeTab) : undefined; return ( -
-
-

-

-
-
+
{demo && } (demo ? setDemoTab(value as Tab) : void setTab(value as Tab))} - className="min-h-0 flex-1 gap-4" + className="min-h-0 flex-1 gap-2" > - - - Traces - - - Investigations - - - +
+
+

+

+ + + Traces + + + Investigations + + +
+
+
+ ({ index: Math.floor((moment(run.start_time).valueOf() - range.startMs) / width), failed: run.error_count > 0, + agent: traceAgentNames(run)[0] ?? "", })); return Array.from({ length: buckets }, (_, i) => { const hits = placed.filter((p) => p.index === i); @@ -40,6 +44,7 @@ export function bucketRuns(runs: readonly TraceSummary[], range: TimeWindow, buc endMs: range.startMs + (i + 1) * width, runs: hits.length, failed: hits.filter((p) => p.failed).length, + agents: hits.filter((p) => !p.failed).map((p) => p.agent), }; }); } @@ -73,44 +78,202 @@ interface TracesTimelineProps { onSelect: (selection: TimeWindow | null) => void; } -function BucketBar({ - bucket, - max, - dimmed, - hovered, -}: { - bucket: Bucket; +const FAILED_RED = "#e5484d"; +const DOT_PITCH = 6; +const FIELD_HEIGHT = FIELD_ROWS * DOT_PITCH; + +function fieldDotColor(lit: boolean, dot: ReturnType[number], grid: string): string { + if (!lit) return grid; + if (dot.kind === "failed") return FAILED_RED; + return dot.kind === "run" ? dot.color : grid; +} + +interface FieldFrame { + buckets: readonly Bucket[]; max: number; - dimmed: boolean; - hovered: boolean; + band: Band | null; + hover: number | null; + progress: number; +} + +function dotRadius(lit: boolean, hovered: boolean): number { + if (!lit) return 0.9; + return hovered ? 2.4 : 1.9; +} + +function drawField(canvas: HTMLCanvasElement, { buckets, max, band, hover, progress }: FieldFrame) { + const context = canvas.getContext("2d"); + const width = canvas.clientWidth; + if (!context || width === 0) return; + const ratio = window.devicePixelRatio || 1; + canvas.width = Math.round(width * ratio); + canvas.height = Math.round(FIELD_HEIGHT * ratio); + context.setTransform(ratio, 0, 0, ratio, 0, 0); + context.clearRect(0, 0, width, FIELD_HEIGHT); + const grid = document.documentElement.classList.contains("dark") ? "rgba(255,255,255,0.08)" : "rgba(15,23,42,0.08)"; + const columnWidth = width / buckets.length; + const pitch = columnWidth / FIELD_COLS; + buckets.forEach((bucket, i) => { + const dots = columnDots(bucket, max, i * 7919 + bucket.runs); + const visible = Math.ceil(dots.length * progress); + const dimmed = band !== null && (i < band.lo || i > band.hi); + dots.forEach((dot, index) => { + const lit = dot.kind !== "grid" && index < visible; + context.globalAlpha = lit && dimmed ? 0.15 : 1; + context.fillStyle = fieldDotColor(lit, dot, grid); + context.beginPath(); + context.arc( + i * columnWidth + pitch * ((index % FIELD_COLS) + 0.5), + FIELD_HEIGHT - DOT_PITCH * (Math.floor(index / FIELD_COLS) + 0.5), + dotRadius(lit, hover === i), + 0, + Math.PI * 2, + ); + context.fill(); + }); + }); + context.globalAlpha = 1; +} + +function useRiseIn(): number { + const [progress, setProgress] = useState(0); + useEffect(() => { + if (typeof window.matchMedia !== "function" || window.matchMedia("(prefers-reduced-motion: reduce)").matches) { + setProgress(1); + return; + } + const start = performance.now(); + let frame = requestAnimationFrame(function rise(now) { + const t = Math.min(1, (now - start) / 700); + setProgress(1 - (1 - t) ** 3); + if (t < 1) frame = requestAnimationFrame(rise); + }); + return () => cancelAnimationFrame(frame); + }, []); + return progress; +} + +function DotField({ + buckets, + max, + band, + hover, +}: { + buckets: readonly Bucket[]; + max: number; + band: Band | null; + hover: number | null; }) { + const canvas = useRef(null); + const progress = useRiseIn(); + useEffect(() => { + const node = canvas.current; + if (!node) return; + const frame: FieldFrame = { buckets, max, band, hover, progress }; + const paint = () => drawField(node, frame); + paint(); + if (typeof ResizeObserver === "undefined") return; + const observer = new ResizeObserver(paint); + observer.observe(node); + return () => observer.disconnect(); + }, [buckets, max, band, hover, progress]); return ( -
- {bucket.runs > 0 && ( -
- {bucket.failed > 0 && ( -
- )} -
+