Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_internal_copy_38013

# Conflicts:
#	litellm/proxy/management_endpoints/model_management_endpoints.py
#	tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py
This commit is contained in:
mateo-berri 2026-09-03 14:12:33 -07:00
commit 6c7d3af053
407 changed files with 72145 additions and 2596 deletions

View file

@ -4,6 +4,11 @@ description: >-
by a job nor listed here, so every entry below is a decision on the record.
test_paths:
- reason: >-
The Rust/Python parity harness is run manually through its local CLI. Recorded replay,
fixture generation, and harness checks are intentionally outside pull request CI
paths:
- tests/rust-python-harness
- reason: >-
What is left of the caching suite in tests/local_testing that runs nowhere. Every job that
globs that directory either deselects it (local_testing_part1 and part2 carry `-k "... and

View file

@ -57,7 +57,7 @@
"limit": 5601
},
"reportMissingTypeArgument": {
"limit": 15288
"limit": 15287
},
"reportMissingTypeStubs": {
"limit": 40
@ -105,10 +105,10 @@
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38324
"limit": 38323
},
"reportUnknownParameterType": {
"limit": 19625
"limit": 19624
},
"reportUnknownVariableType": {
"limit": 29861

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.63"
version = "0.1.64"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.63"
version = "0.1.64"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.92"
version = "0.4.93"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.92"
version = "0.4.93"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -27,7 +27,7 @@ rand = "0.8"
reqwest = { version = "0.12", default-features = false, features = ["blocking", "json", "rustls-tls", "http2", "stream"] }
rstest = "0.26.1"
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
serde_json = { version = "1.0", features = ["float_roundtrip"] }
sha2 = "0.10"
subtle = "2"
thiserror = "2.0"

View file

@ -14,7 +14,7 @@ const AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT_ENV: &str = "AZURE_DOCUMENT_INTELLIGE
const AZURE_DOCUMENT_INTELLIGENCE_API_VERSION: &str = "2024-11-30";
const AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI: i64 = 96;
const AZURE_DOCUMENT_INTELLIGENCE_SUPPORTED_OCR_PARAMS: &[&str] = &["pages"];
const AZURE_DOCUMENT_INTELLIGENCE_SUPPORTED_OCR_PARAMS: &[&str] = &["pages", "features"];
pub struct AzureAiOcrConfig;
pub struct AzureDocumentIntelligenceOcrConfig;
@ -192,6 +192,46 @@ fn normalize_pages_param(pages: &Value) -> Result<Option<String>, Error> {
}
}
fn feature_token_is_valid(token: &str) -> bool {
let Some((first, rest)) = token.as_bytes().split_first() else {
return false;
};
first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric)
}
fn invalid_features_error(features: &Value) -> Error {
Error::InvalidRequest(format!(
"Invalid `features` for Azure Document Intelligence: {features:?}. Expected a list of feature names or a comma-separated string like 'keyValuePairs' or 'keyValuePairs,languages'."
))
}
fn normalize_features_param(features: &Value) -> Result<Option<String>, Error> {
let normalized = match features {
Value::String(value) => value
.split(',')
.map(str::trim)
.collect::<Vec<_>>()
.join(","),
Value::Array(values) if values.is_empty() => return Ok(None),
Value::Array(values) => values
.iter()
.map(Value::as_str)
.collect::<Option<Vec<_>>>()
.ok_or_else(|| invalid_features_error(features))?
.into_iter()
.map(str::trim)
.collect::<Vec<_>>()
.join(","),
_ => return Err(invalid_features_error(features)),
};
if normalized.split(',').all(feature_token_is_valid) {
Ok(Some(normalized))
} else {
Err(invalid_features_error(features))
}
}
pub fn complete_document_intelligence_url(
api_base: Option<&str>,
model: &str,
@ -213,6 +253,13 @@ pub fn complete_document_intelligence_url(
url.push_str(&normalized);
}
if let Some(features) = optional_params.get("features")
&& let Some(normalized) = normalize_features_param(features)?
{
url.push_str("&features=");
url.push_str(&normalized);
}
Ok(url)
}
@ -475,6 +522,103 @@ mod tests {
);
}
#[test]
fn document_intelligence_url_normalizes_features() {
let params = serde_json::Map::from_iter([(
"features".to_string(),
json!("keyValuePairs, languages"),
)]);
let url = complete_document_intelligence_url(
Some("https://example.cognitiveservices.azure.com"),
"prebuilt-layout",
&params,
&|_| None,
)
.expect("url builds");
assert_eq!(
url,
"https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30&features=keyValuePairs,languages"
);
}
#[test]
fn document_intelligence_url_combines_pages_and_feature_list() {
let params = serde_json::Map::from_iter([
("pages".to_string(), json!([0, 1, 2])),
(
"features".to_string(),
json!([" keyValuePairs ", "languages"]),
),
]);
let url = complete_document_intelligence_url(
Some("https://example.cognitiveservices.azure.com"),
"prebuilt-layout",
&params,
&|_| None,
)
.expect("url builds");
assert_eq!(
url,
"https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30&pages=1,2,3&features=keyValuePairs,languages"
);
}
#[test]
fn document_intelligence_url_omits_empty_feature_list() {
let params = serde_json::Map::from_iter([("features".to_string(), json!([]))]);
let url = complete_document_intelligence_url(
Some("https://example.cognitiveservices.azure.com"),
"prebuilt-layout",
&params,
&|_| None,
)
.expect("url builds");
assert_eq!(
url,
"https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30"
);
}
#[test]
fn document_intelligence_url_rejects_invalid_features() {
for features in [
json!("keyValuePairs&pages=9"),
json!(""),
json!(["keyValuePairs", 1]),
json!({"feature": "keyValuePairs"}),
] {
let params = serde_json::Map::from_iter([("features".to_string(), features.clone())]);
let error = complete_document_intelligence_url(
Some("https://example.cognitiveservices.azure.com"),
"prebuilt-layout",
&params,
&|_| None,
)
.expect_err("invalid features must fail");
assert!(
matches!(error, Error::InvalidRequest(message) if message.contains("Invalid `features`")),
"features={features:?}"
);
}
}
#[test]
fn document_intelligence_maps_features() {
let params = Map::from_iter([
("features".to_string(), json!(["keyValuePairs"])),
("unsupported".to_string(), json!(true)),
]);
assert_eq!(
AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG.map_ocr_params(&params),
Map::from_iter([("features".to_string(), json!(["keyValuePairs"]))])
);
}
#[test]
fn document_intelligence_request_uses_base64_source_for_data_uri() {
let body = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG

View file

@ -59,3 +59,41 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())
}
pub(crate) fn ocr_error_to_pyerr(err: Error) -> PyErr {
match err {
Error::MissingField("document_url" | "image_url") => {
PyValueError::new_err("Document URL is required")
}
Error::Http { status, body } => RustUpstreamError::new_err((status, body)),
other => core_error_to_pyerr(other),
}
}
#[cfg(test)]
mod ocr_error_tests {
use super::*;
#[test]
fn ocr_errors_preserve_python_validation_and_provider_details() {
Python::initialize();
Python::attach(|py| {
for field in ["document_url", "image_url"] {
let mapped = ocr_error_to_pyerr(Error::MissingField(field));
assert!(mapped.is_instance_of::<PyValueError>(py));
assert_eq!(mapped.value(py).to_string(), "Document URL is required");
}
let mapped = ocr_error_to_pyerr(Error::Http {
status: 429,
body: r#"{"message":"rate limited"}"#.to_string(),
});
assert!(mapped.is_instance_of::<RustUpstreamError>(py));
let args: (u16, String) = mapped
.value(py)
.getattr("args")
.and_then(|args| args.extract())
.expect("OCR failures retain status and unprefixed provider message");
assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string()));
});
}
}

View file

@ -5,7 +5,7 @@ use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
use pyo3::prelude::*;
use serde_json::Value;
use crate::errors::core_error_to_pyerr;
use crate::errors::ocr_error_to_pyerr;
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
fn prepare_ocr(
@ -69,5 +69,5 @@ bridge_route! {
timeout_seconds: Option<f64>,
},
prepare = prepare_ocr,
errors = core_error_to_pyerr,
errors = ocr_error_to_pyerr,
}

View file

@ -932,7 +932,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return OpenAiResponsesToChatCompletionStreamIterator(streaming_response, sync_stream, json_mode)
def _convert_content_str_to_input_text(self, content: str, role: str) -> dict[str, object]:
if role == "user" or role == "system" or role == "tool":
if role in ("user", "system", "developer", "tool"):
return {"type": "input_text", "text": content}
else:
return {"type": "output_text", "text": content}

View file

@ -7,6 +7,7 @@ from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_in_ran
DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"))
AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))
AZURE_OPENAI_AUDIO_PROVIDERS: Final = frozenset({"azure", "azure_ai"})
ROUTER_MAX_FALLBACKS: Final = int(os.getenv("ROUTER_MAX_FALLBACKS", 5))
ROUTER_FALLBACK_ERROR_DETAIL_MAX_CHARS: Final = 2000
RUNTIME_UPDATABLE_ROUTER_SETTINGS: Final[frozenset[str]] = frozenset(
@ -39,6 +40,7 @@ ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
"router_general_settings",
"ignore_invalid_deployments",
"fallback_access_check",
"heuristic_v2_router_limit",
}
)
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
@ -1450,6 +1452,7 @@ SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY: Final = "_session_deployment_affin
CONSUMED_REQUEST_TAGS_METADATA_KEY: Final = "_consumed_request_tags"
INTERNAL_CALL_ORIGIN_METADATA_KEY: Final = "internal_call_origin"
SESSION_ID_GENERATED_METADATA_KEY: Final = "litellm_session_id_generated"
SESSION_ID_OMITTED_METADATA_KEY: Final = "litellm_session_id_omitted"
LITELLM_TRUNCATED_PAYLOAD_FIELD: Final = "litellm_truncated"
LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE: Final = (
"Truncation is a DB storage safeguard. "

View file

@ -5,6 +5,8 @@ Helper utilities for tracking the cost of built-in tools.
from collections.abc import Mapping
from typing import Final, Literal
from pydantic import ValidationError
import litellm
from litellm.constants import OPENAI_FILE_SEARCH_COST_PER_1K_CALLS
from litellm.litellm_core_utils.llm_cost_calc.utils import (
@ -13,6 +15,7 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import (
from litellm.types.llms.openai import (
FileSearchTool,
ResponsesAPIResponse,
ResponsesToolUsage,
WebSearchOptions,
)
from litellm.types.utils import (
@ -32,6 +35,17 @@ def _output_item_type(output_item: object) -> str | None:
return item_type if isinstance(item_type, str) else None
def _reported_web_search_requests(response_object: ResponsesAPIResponse) -> int | None:
tool_usage: Final = getattr(response_object, "tool_usage", None)
if tool_usage is None:
return None
try:
web_search: Final = ResponsesToolUsage.model_validate(tool_usage).web_search
except ValidationError:
return None
return None if web_search is None else web_search.num_requests
def _usage_reports_server_side_web_search_calls(usage: Usage) -> bool:
details: Final = getattr(usage, "server_side_tool_usage_details", None)
if not isinstance(details, Mapping):
@ -182,15 +196,19 @@ class StandardBuiltInToolCostTracking:
Providers that report a request count in usage (gemini, anthropic, xai, vertex) are handled by
get_cost_for_web_search_request and never reach here. This path prices per call, so it must count
the web_search_call items. Chat-completions responses only expose url_citation annotations with no
count, so they floor to a single billable search.
the web_search_call items, unless the response reports the billable count itself
(Bedrock's tool_usage.web_search.num_requests, which excludes open_page fetches). Chat-completions
responses only expose url_citation annotations with no count, so they floor to a single billable search.
"""
if isinstance(response_object, ResponsesAPIResponse):
count = sum(
1 for output_item in response_object.output if _output_item_type(output_item) == "web_search_call"
)
return max(count, 1)
return 1
if not isinstance(response_object, ResponsesAPIResponse):
return 1
reported: Final = _reported_web_search_requests(response_object)
if reported is not None:
return reported
count: Final = sum(
1 for output_item in response_object.output if _output_item_type(output_item) == "web_search_call"
)
return max(count, 1)
@staticmethod
def _handle_file_search_cost(

View file

@ -428,7 +428,7 @@ def _coerce_off_peak_rate(value: object, default: float) -> float:
return default
def _apply_off_peak_pricing(
def apply_off_peak_pricing(
model_info: ModelInfo,
current_time: datetime | None,
prompt_base_cost: float,
@ -462,7 +462,7 @@ def _apply_off_peak_to_base_costs(
has no field for them.
"""
prompt, completion, cache_creation, cache_creation_above_1hr, cache_read = base_costs
off_peak_prompt, off_peak_completion, off_peak_cache_read = _apply_off_peak_pricing(
off_peak_prompt, off_peak_completion, off_peak_cache_read = apply_off_peak_pricing(
model_info, current_time, prompt, completion, cache_read
)
return (off_peak_prompt, off_peak_completion, cache_creation, cache_creation_above_1hr, off_peak_cache_read)

View file

@ -1554,6 +1554,22 @@ def with_prompt_cache_breakpoint(target: _MarkedT, marker: object) -> _MarkedT:
return cast(_MarkedT, marked) # cast-ok: same block shape as the input plus the marker key
LITELLM_INTERNAL_MESSAGE_FIELDS: Final = frozenset({"thinking_blocks", "reasoning_content", "provider_specific_fields"})
def strip_litellm_internal_message_fields(message: AllMessageValues) -> AllMessageValues:
"""Drop the fields litellm attaches to assistant messages (e.g. when translating Anthropic thinking
blocks) that OpenAI-compatible endpoints with strict schemas reject as extra inputs."""
if LITELLM_INTERNAL_MESSAGE_FIELDS.isdisjoint(message):
return message
return cast( # cast-ok: same TypedDict minus internal keys
AllMessageValues,
{ # mutable-ok: provider transforms mutate message dicts in place downstream
key: value for key, value in message.items() if key not in LITELLM_INTERNAL_MESSAGE_FIELDS
},
)
def filter_value_from_dict(dictionary: dict, key: str, depth: int = 0) -> Any:
"""
Filters a value from a dictionary

View file

@ -2337,6 +2337,9 @@ class CustomStreamWrapper:
else:
self.sent_last_chunk = True
processed_chunk: Final = self.finish_reason_handler()
if self.stream_options is None:
usage: Final = calculate_total_usage(chunks=self.chunks)
processed_chunk._hidden_params["usage"] = usage # pyright: ignore[reportPrivateUsage] # sync parity
# see sync __next__'s sibling branch: deliberately do NOT restore
# here - this chunk is still this call's own data, and restoring
# before returning it would corrupt the caller's own log

View file

@ -15,6 +15,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
filter_value_from_dict,
)
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.azure_ai.common_utils import is_foundry_model_inference_base
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
from litellm.llms.openai.openai import OpenAIConfig
@ -207,20 +208,18 @@ class AzureAIStudioConfig(OpenAIConfig):
message["content"] = texts
return stripped_messages
def _is_azure_openai_model(self, model: str, api_base: str | None) -> bool:
try:
if "/" in model:
model = model.split("/", 1)[1]
if (
model in litellm.open_ai_chat_completion_models
or model in litellm.open_ai_text_completion_models
or model in litellm.open_ai_embedding_models
):
return True
def _is_foundry_model_inference_base(self, api_base: str) -> bool:
return is_foundry_model_inference_base(api_base)
except Exception:
def _is_azure_openai_model(self, model: str, api_base: str | None) -> bool:
if api_base is None or self._is_foundry_model_inference_base(api_base):
return False
return False
stripped_model: Final = model.split("/", 1)[1] if "/" in model else model
return (
stripped_model in litellm.open_ai_chat_completion_models
or stripped_model in litellm.open_ai_text_completion_models
or stripped_model in litellm.open_ai_embedding_models
)
def _get_openai_compatible_provider_info(
self,

View file

@ -1,5 +1,6 @@
from collections.abc import Mapping
from typing import Final, Literal
from urllib.parse import urlparse
import litellm
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
@ -10,6 +11,14 @@ from litellm.types.router import GenericLiteLLMParams
AzureAIApiKeyHeader = Literal["Authorization", "api-key", "Api-Key", "Ocp-Apim-Subscription-Key"]
def is_foundry_model_inference_base(api_base: str) -> bool:
parsed: Final = urlparse(api_base)
host: Final = parsed.hostname
if host is None or not host.endswith(".services.ai.azure.com"):
return False
return "/openai/deployments" not in parsed.path
def get_azure_ai_entra_token(litellm_params: Mapping[str, object] | None = None) -> str | None:
"""
Resolve an Entra ID / OAuth access token for an Azure AI Foundry deployment.

View file

@ -1,8 +1,10 @@
from typing import Final
from urllib.parse import urlsplit, urlunsplit
from openai import OpenAI
import litellm
from litellm.llms.azure_ai.common_utils import is_foundry_model_inference_base
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -16,6 +18,16 @@ from litellm.utils import convert_to_model_response_object
from .cohere_transformation import AzureAICohereConfig
def _foundry_models_route_base(api_base: str | None) -> str | None:
if api_base is None or not is_foundry_model_inference_base(api_base):
return api_base
parts: Final = urlsplit(api_base)
path: Final = parts.path.rstrip("/")
if path.endswith("/models"):
return api_base
return urlunsplit((parts.scheme, parts.netloc, f"{path}/models", parts.query, parts.fragment))
class AzureAIEmbedding(OpenAIChatCompletion):
def _process_response(
self,
@ -214,6 +226,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
assemble result in-order, and return
"""
resolved_api_base: Final = _foundry_models_route_base(api_base)
if aembedding is True:
return self.async_embedding(
model,
@ -223,7 +236,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
model_response,
optional_params,
api_key,
api_base,
resolved_api_base,
client,
)
@ -245,7 +258,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
model_response=model_response,
optional_params=optional_params,
api_key=api_key,
api_base=api_base,
api_base=resolved_api_base,
client=client,
)
@ -262,7 +275,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
model_response,
optional_params,
api_key,
api_base,
resolved_api_base,
client=(client if client is not None and isinstance(client, OpenAI) else None),
aembedding=aembedding,
shared_session=shared_session,

View file

@ -48,7 +48,9 @@ _BASE_SUFFIXES_TO_STRIP: Final = (
)
# Per Bedrock Mantle Responses API validation errors.
_BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES = frozenset({"function", "mcp", "custom", "namespace", "tool_search"})
_BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES: Final = frozenset(
{"function", "mcp", "custom", "namespace", "tool_search", "web_search"}
)
_BEDROCK_MANTLE_SUPPORTED_SERVICE_TIERS: Final = frozenset({"auto", "default"})

View file

@ -113,9 +113,12 @@ from litellm.types.containers.main import (
)
from litellm.types.files import StreamingMediaUploadConfig, TwoStepFileUploadConfig
from litellm.types.integrations.custom_logger import (
NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES,
AgenticLoopPlan,
AgenticLoopRequestPatch,
AgenticLoopSafetyError,
converted_stream_requested,
is_interception_internal_key,
)
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
@ -2823,6 +2826,7 @@ class BaseLLMHTTPHandler:
)
if self._has_agentic_completion_hook(logging_obj):
agentic_kwargs: Final = dict(litellm_params) # mutable-ok: agentic hooks mutate kwargs in place
final_response: Final = run_async_function(
self._call_agentic_completion_hooks,
response=initial_response,
@ -2833,10 +2837,19 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
stream=False,
custom_llm_provider=custom_llm_provider,
kwargs=dict(litellm_params),
kwargs=agentic_kwargs,
api_surface="responses",
)
return final_response if final_response is not None else initial_response
result: Final = final_response if final_response is not None else initial_response
if converted_stream_requested(agentic_kwargs) and not agentic_kwargs.get("_agentic_loop_depth"):
return self._wrap_responses_response_as_fake_stream(
result=result,
model=model,
responses_api_provider_config=responses_api_provider_config,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
)
return result
return initial_response
@ -3002,6 +3015,7 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
)
agentic_kwargs: Final = dict(litellm_params) # mutable-ok: agentic hooks mutate kwargs in place
final_response: Final = await self._call_agentic_completion_hooks(
response=initial_response,
model=model,
@ -3011,15 +3025,12 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
stream=False,
custom_llm_provider=custom_llm_provider,
kwargs=dict(litellm_params),
kwargs=agentic_kwargs,
api_surface="responses",
)
result: Final = final_response if final_response is not None else initial_response
interception_converted_stream: Final = litellm_params.get(
"_code_interpreter_interception_converted_stream"
) or litellm_params.get("_websearch_interception_converted_stream")
if interception_converted_stream and not litellm_params.get("_agentic_loop_depth"):
if converted_stream_requested(agentic_kwargs) and not agentic_kwargs.get("_agentic_loop_depth"):
return self._wrap_responses_response_as_fake_stream(
result=result,
model=model,
@ -5583,8 +5594,7 @@ class BaseLLMHTTPHandler:
kwargs_for_followup: Final = {
k: v
for k, v in kwargs.items()
if not k.startswith("_websearch_interception")
and not k.startswith("_compression_interception")
if not is_interception_internal_key(k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES)
and k != "_code_interpreter_interception_converted_stream"
and k not in internal_keys
and k not in optional_params

View file

@ -7,11 +7,13 @@ cached, cache-creation, output, reasoning) is billed at that one tier's rate.
See https://help.aliyun.com/zh/model-studio/billing-for-model-studio
"""
from dataclasses import dataclass
from dataclasses import dataclass, replace
from datetime import datetime
from typing import Final
from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import select_tier_for_input, tier_rate
from litellm.litellm_core_utils.llm_cost_calc.utils import (
apply_off_peak_pricing,
parse_completion_tokens_details,
parse_prompt_tokens_details,
)
@ -32,6 +34,19 @@ class TokenBreakdown:
return self.text_tokens + self.cached_tokens + self.cache_creation_tokens
@dataclass(frozen=True, slots=True)
class TokenRates:
input_rate: float
cache_read_rate: float
cache_creation_rate: float
output_rate: float
reasoning_rate: float | None
@property
def billed_reasoning_rate(self) -> float:
return self.output_rate if self.reasoning_rate is None else self.reasoning_rate
def _extract_token_breakdown(usage: Usage) -> TokenBreakdown:
prompt_details: Final = parse_prompt_tokens_details(usage)
cached_tokens: Final = prompt_details["cache_hit_tokens"]
@ -57,69 +72,75 @@ def _flat_rate(model_info: ModelInfo, cost_key: str, fallback_cost_key: str) ->
return float(value)
def _calculate_prompt_cost(
breakdown: TokenBreakdown,
model_info: ModelInfo,
tier: dict | None,
) -> float:
if tier is not None:
return (
(breakdown.text_tokens * tier_rate(tier, "input_cost_per_token"))
+ (breakdown.cached_tokens * tier_rate(tier, "cache_read_input_token_cost", "input_cost_per_token"))
+ (
breakdown.cache_creation_tokens
* tier_rate(tier, "cache_creation_input_token_cost", "input_cost_per_token")
)
)
input_cost: Final = float(model_info.get("input_cost_per_token") or 0.0)
cache_read_cost: Final = _flat_rate(model_info, "cache_read_input_token_cost", "input_cost_per_token")
cache_creation_cost: Final = _flat_rate(model_info, "cache_creation_input_token_cost", "input_cost_per_token")
return (
(breakdown.text_tokens * input_cost)
+ (breakdown.cached_tokens * cache_read_cost)
+ (breakdown.cache_creation_tokens * cache_creation_cost)
def _flat_rates(model_info: ModelInfo) -> TokenRates:
reasoning_rate: Final = model_info.get("output_cost_per_reasoning_token")
return TokenRates(
input_rate=float(model_info.get("input_cost_per_token") or 0.0),
cache_read_rate=_flat_rate(model_info, "cache_read_input_token_cost", "input_cost_per_token"),
cache_creation_rate=_flat_rate(model_info, "cache_creation_input_token_cost", "input_cost_per_token"),
output_rate=float(model_info.get("output_cost_per_token") or 0.0),
reasoning_rate=None if reasoning_rate is None else float(reasoning_rate),
)
def _calculate_completion_cost(
breakdown: TokenBreakdown,
model_info: ModelInfo,
tier: dict | None,
) -> float:
def _tier_rates(model_info: ModelInfo, tier: dict) -> TokenRates:
# A tier that declares output rates keeps the request on them, all-or-nothing. A tier table
# spelling out only input rates would serve every completion for free, so there the model's
# own output rates stand in
tier_declares_output: Final = tier is not None and "output_cost_per_token" in tier
output_cost: Final = (
tier_rate(tier, "output_cost_per_token")
if tier_declares_output
else float(model_info.get("output_cost_per_token") or 0.0)
)
tier_declares_reasoning: Final = tier is not None and "output_cost_per_reasoning_token" in tier
model_reasoning_rate: Final = None if tier_declares_output else model_info.get("output_cost_per_reasoning_token")
reasoning_cost: Final = (
tier_rate(tier, "output_cost_per_reasoning_token", "output_cost_per_token")
if tier_declares_reasoning
else float(model_reasoning_rate)
if model_reasoning_rate is not None
else output_cost
flat_rates: Final = _flat_rates(model_info)
tier_declares_output: Final = "output_cost_per_token" in tier
tier_declares_reasoning: Final = "output_cost_per_reasoning_token" in tier
return TokenRates(
input_rate=tier_rate(tier, "input_cost_per_token"),
cache_read_rate=tier_rate(tier, "cache_read_input_token_cost", "input_cost_per_token"),
cache_creation_rate=tier_rate(tier, "cache_creation_input_token_cost", "input_cost_per_token"),
output_rate=tier_rate(tier, "output_cost_per_token") if tier_declares_output else flat_rates.output_rate,
reasoning_rate=(
tier_rate(tier, "output_cost_per_reasoning_token")
if tier_declares_reasoning
else None
if tier_declares_output
else flat_rates.reasoning_rate
),
)
return (breakdown.completion_tokens * output_cost) + (breakdown.reasoning_tokens * reasoning_cost)
def _off_peak_rates(model_info: ModelInfo, current_time: datetime | None, rates: TokenRates) -> TokenRates:
input_rate, output_rate, cache_read_rate = apply_off_peak_pricing(
model_info, current_time, rates.input_rate, rates.output_rate, rates.cache_read_rate
)
return replace(rates, input_rate=input_rate, output_rate=output_rate, cache_read_rate=cache_read_rate)
def cost_per_token(model: str, usage: Usage, custom_llm_provider: str = "dashscope") -> tuple[float, float]:
def _bill(breakdown: TokenBreakdown, rates: TokenRates) -> tuple[float, float]:
prompt_cost: Final = (
(breakdown.text_tokens * rates.input_rate)
+ (breakdown.cached_tokens * rates.cache_read_rate)
+ (breakdown.cache_creation_tokens * rates.cache_creation_rate)
)
completion_cost: Final = (breakdown.completion_tokens * rates.output_rate) + (
breakdown.reasoning_tokens * rates.billed_reasoning_rate
)
return prompt_cost, completion_cost
def cost_per_token(
model: str,
usage: Usage,
custom_llm_provider: str = "dashscope",
current_time: datetime | None = None,
) -> tuple[float, float]:
"""
Calculate cost per token for Dashscope models.
Supports both tiered and flat pricing with cached and reasoning tokens.
Supports both tiered and flat pricing with cached and reasoning tokens, and swaps in the
model's off_peak_pricing rates while one of its windows is open.
Args:
model: Model name without provider prefix
usage: LiteLLM Usage block
custom_llm_provider: The provider id the request resolved to; dashscope or one of its brand aliases
current_time: The moment the request is billed at; defaults to now, UTC
Returns:
Tuple[float, float] - (prompt_cost_in_usd, completion_cost_in_usd)
@ -133,8 +154,7 @@ def cost_per_token(model: str, usage: Usage, custom_llm_provider: str = "dashsco
if tiered_pricing
else None
)
standard_rates: Final = _flat_rates(model_info) if tier is None else _tier_rates(model_info, tier)
rates: Final = _off_peak_rates(model_info, current_time, standard_rates)
prompt_cost: Final = _calculate_prompt_cost(breakdown=breakdown, model_info=model_info, tier=tier)
completion_cost: Final = _calculate_completion_cost(breakdown=breakdown, model_info=model_info, tier=tier)
return prompt_cost, completion_cost
return _bill(breakdown, rates)

View file

@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to Databricks' `/chat/completion
"""
import os
from collections.abc import AsyncIterator, Coroutine, Iterator
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, cast, overload
import httpx
@ -15,6 +15,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
_should_convert_tool_call_to_json_mode,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import (
strip_litellm_internal_message_fields,
strip_name_from_message,
)
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
@ -55,6 +56,14 @@ from ...openai_like.chat.transformation import OpenAILikeChatConfig
from ..common_utils import DatabricksBase, DatabricksException
def _is_bare_assistant_message(message_dict: Mapping[str, object]) -> bool:
"""Databricks rejects assistant messages with neither content nor tool calls, e.g. a replayed
thinking-only turn once its `thinking_blocks` are stripped."""
return message_dict.get("role") == "assistant" and not any(
message_dict.get(key) for key in ("content", "tool_calls", "function_call")
)
def _sanitize_empty_content(message_dict: dict[str, Any]) -> None:
"""
Remove or filter content so empty text blocks are not sent.
@ -423,6 +432,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
"""
Databricks does not support:
- 'name' in user message.
- litellm's internal `thinking_blocks` / `reasoning_content` on assistant messages.
"""
new_messages = []
for idx, message in enumerate(messages):
@ -431,10 +441,13 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
else:
_message = message
_message = strip_name_from_message(_message, allowed_name_roles=["user"])
_message = strip_litellm_internal_message_fields(_message)
# Move message-level cache_control into a content block when content is a string.
if "cache_control" in _message and isinstance(_message.get("content"), str):
_message = self._move_cache_control_into_string_content_block(_message)
_sanitize_empty_content(cast(dict[str, Any], _message))
if _is_bare_assistant_message(_message):
continue
new_messages.append(_message)
if "claude" not in model:

View file

@ -11,6 +11,7 @@ import time
import uuid
from collections.abc import AsyncIterator, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Optional
from urllib.parse import urlsplit
import httpx
import openai
@ -43,6 +44,14 @@ _OPENAI_INIT_PARAMS: Final[tuple[str, ...]] = _get_client_init_params(OpenAI)
_AZURE_OPENAI_INIT_PARAMS: Final[tuple[str, ...]] = _get_client_init_params(AzureOpenAI)
_OPENAI_API_HOST: Final[str] = "api.openai.com"
def is_openai_backed_api_base(api_base: str) -> bool:
hostname: Final = urlsplit(api_base).hostname
return hostname is not None and (hostname == _OPENAI_API_HOST or hostname.endswith(f".{_OPENAI_API_HOST}"))
class OpenAIError(BaseLLMException):
def __init__(
self,

View file

@ -82,8 +82,8 @@ class GPTImageGenerationConfig(BaseImageGenerationConfig):
)
# set optional params
image_response.size = optional_params.get("size", "1024x1024") # default is always 1024x1024
image_response.quality = optional_params.get("quality", "high") # always hd for dall-e-3
image_response.output_format = optional_params.get("response_format", "png") # always png for dall-e-3
image_response.size = image_response.size or optional_params.get("size", "1024x1024")
image_response.quality = image_response.quality or optional_params.get("quality", "high")
image_response.output_format = image_response.output_format or optional_params.get("output_format", "png")
return image_response

View file

@ -2,7 +2,6 @@ import time
import types
from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Iterator, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
from urllib.parse import urlparse
import httpx
@ -55,6 +54,7 @@ from .common_utils import (
OpenAIError,
build_output_token_limit_response,
drop_params_from_unprocessable_entity_error,
is_openai_backed_api_base,
is_output_token_limit_error,
)
from .workload_identity import resolve_openai_workload_identity_config
@ -1190,10 +1190,8 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
"""
if stream_options is not None:
return {"stream_options": stream_options}
else:
# by default litellm will include usage for openai endpoints
if api_base is None or urlparse(api_base).hostname == "api.openai.com":
return {"stream_options": {"include_usage": True}}
if api_base is None or is_openai_backed_api_base(api_base):
return {"stream_options": {"include_usage": True}}
return {}
# Embedding

View file

@ -33,8 +33,9 @@ import time
import uuid
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from itertools import accumulate
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Union, cast
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Union, cast
from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall
from pydantic import BaseModel, TypeAdapter
@ -42,6 +43,7 @@ from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
OpenAiResponsesToChatCompletionStreamIterator,
)
from litellm.llms.base_llm.guardrail_translation.base_translation import (
@ -74,6 +76,7 @@ from litellm.types.llms.openai import (
OutputTextDoneEvent,
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
ResponsesAPIStreamingResponse,
@ -115,6 +118,199 @@ class ResponsesStreamChunk(TypedDict, total=False):
content_index: ReadOnly[int]
_PATCHABLE_ITEM_FIELDS: Final[Mapping[str, str]] = MappingProxyType(
{"function_call_output": "output", "message": "content"}
)
_EMPTY_RESPONSES_REQUEST: Final[ResponsesAPIOptionalRequestParams] = {}
def _item_rewrite_field(item: Mapping[str, object]) -> str | None:
item_type: Final = item.get("type")
if item_type is None:
return "content" if "content" in item else None
if not isinstance(item_type, str):
return None
return _PATCHABLE_ITEM_FIELDS.get(item_type)
def _rewritten_input_item(item: Mapping[str, object], rewritten: object) -> Mapping[str, object] | None:
field: Final = _item_rewrite_field(item)
if field is None or not isinstance(rewritten, Mapping):
return None
rewritten_content: Final = rewritten.get("content")
if isinstance(item.get(field), str) and isinstance(rewritten_content, str):
return {**item, field: rewritten_content} # mutable-ok: request input items must stay JSON-plain dicts
rewritten_row: Final = cast("AllMessageValues", rewritten) # cast-ok: guardrails hand back chat-shaped rows
converted_items, _ = LiteLLMResponsesTransformationHandler().convert_chat_completion_messages_to_responses_api(
[rewritten_row] # mutable-ok: converter signature takes a list
)
if len(converted_items) != 1 or not isinstance(converted_items[0], Mapping):
return None
first_converted: Final = cast("Mapping[str, object]", converted_items[0]) # cast-ok: isinstance-checked above
converted_value: Final = first_converted.get(field)
if converted_value is None:
return None
return {**item, field: converted_value} # mutable-ok: request input items must stay JSON-plain dicts
def _is_function_call_item(item: object) -> bool:
return isinstance(item, Mapping) and item.get("type") in ("function_call", "custom_tool_call")
def _last_message_role(messages: Sequence[object]) -> str | None:
if not messages:
return None
last: Final = messages[-1]
role: Final = last.get("role") if isinstance(last, Mapping) else getattr(last, "role", None)
return role if isinstance(role, str) else None
def _provenance_unit_bounds(
raw_input: Sequence[object],
solo_conversions: Sequence[Sequence[object]],
) -> tuple[tuple[int, int], ...]:
trailing_roles: Final = tuple(
accumulate(
(_last_message_role(messages) for messages in solo_conversions),
lambda previous, current: current if current is not None else previous,
)
)
start_indexes: Final = tuple(
index
for index in range(len(raw_input))
if index == 0 or not (_is_function_call_item(raw_input[index]) and trailing_roles[index - 1] == "assistant")
)
return tuple(zip(start_indexes, (*start_indexes[1:], len(raw_input))))
def _input_item_provenance(
raw_input: Sequence[object],
expected_messages: Sequence[object],
) -> tuple[Mapping[int, int], frozenset[int]] | None:
if not all(isinstance(item, Mapping) for item in raw_input):
return None
solo_conversions: Final = tuple(
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=cast("ResponseInputParam", [item]), # cast-ok: items checked as Mappings above
responses_api_request=_EMPTY_RESPONSES_REQUEST,
)
for item in raw_input
)
full_conversion: Final = tuple(
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=cast("ResponseInputParam", list(raw_input)), # cast-ok: items checked as Mappings above
responses_api_request=_EMPTY_RESPONSES_REQUEST,
)
)
if full_conversion != tuple(expected_messages):
return None
units: Final = _provenance_unit_bounds(raw_input, solo_conversions)
unit_messages: Final = tuple(
tuple(solo_conversions[start])
if end - start == 1
else tuple(
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
input=cast("ResponseInputParam", list(raw_input[start:end])), # cast-ok: checked as Mappings above
responses_api_request=_EMPTY_RESPONSES_REQUEST,
)
)
for start, end in units
)
if tuple(message for messages in unit_messages for message in messages) != full_conversion:
return None
boundaries: Final = tuple(accumulate((len(messages) for messages in unit_messages), initial=0))
item_for_message: Final = MappingProxyType(
{
message_index: start
for unit_index, (start, end) in enumerate(units)
if end - start == 1
for message_index in range(boundaries[unit_index], boundaries[unit_index + 1])
}
)
tainted: Final = frozenset(
message_index
for unit_index, (start, end) in enumerate(units)
if end - start > 1
for message_index in range(boundaries[unit_index], boundaries[unit_index + 1])
)
return item_for_message, tainted
class _RequestFields(NamedTuple):
input: tuple[object, ...]
instructions: str | None
class _ExtractedInputs(NamedTuple):
inputs: GenericGuardrailAPIInputs
task_mappings: tuple[tuple[int, int | None], ...]
def _patched_request_fields(
raw_input: object,
instructions: object,
original_messages: Sequence[object],
structured_messages: Sequence[object],
) -> _RequestFields | None:
if not isinstance(raw_input, list) or len(original_messages) != len(structured_messages):
return None
offset: Final = 1 if instructions else 0
provenance: Final = _input_item_provenance(raw_input, tuple(original_messages)[offset:])
if provenance is None:
return None
item_for_message, tainted = provenance
changed: Final = tuple(
(index, rewritten)
for index, (original, rewritten) in enumerate(zip(original_messages, structured_messages))
if original != rewritten
)
instruction_rewrites: Final = tuple(rewritten for index, rewritten in changed if index < offset)
rewritten_instructions: Final = (
instruction_rewrites[0].get("content")
if instruction_rewrites and isinstance(instruction_rewrites[0], Mapping)
else instructions
)
instructions_value: Final = rewritten_instructions if isinstance(rewritten_instructions, str) else None
if rewritten_instructions is not None and instructions_value is None:
return None
body_changes: Final = tuple((index - offset, rewritten) for index, rewritten in changed if index >= offset)
if any(message_index in tainted or message_index not in item_for_message for message_index, _ in body_changes):
return None
replacements: Final = MappingProxyType(
{
item_for_message[message_index]: _rewritten_input_item(
cast("Mapping[str, object]", raw_input[item_for_message[message_index]]), # cast-ok: checked Mappings
rewritten,
)
for message_index, rewritten in body_changes
}
)
if len(replacements) != len(body_changes) or any(item is None for item in replacements.values()):
return None
return _RequestFields(
input=tuple(replacements.get(index, item) for index, item in enumerate(raw_input)),
instructions=instructions_value,
)
def _patch_or_convert_request_fields(
raw_input: object,
instructions: object,
original_messages: Sequence[object],
structured_messages: Sequence[AllMessageValues],
) -> _RequestFields | None:
if not isinstance(structured_messages, list):
return None
patched: Final = _patched_request_fields(raw_input, instructions, original_messages, structured_messages)
if patched is not None:
return patched
input_items, converted_instructions = (
LiteLLMResponsesTransformationHandler().convert_chat_completion_messages_to_responses_api(structured_messages)
)
return _RequestFields(input=tuple(input_items), instructions=converted_instructions)
def _next_stream_sequence_number(responses_so_far: Sequence[Any] | None) -> int:
sequence_numbers: Final = (
item.get("sequence_number") if isinstance(item, dict) else getattr(item, "sequence_number", None)
@ -162,9 +358,8 @@ class OpenAIResponsesHandler(BaseTranslation):
Handles both string input and list of message objects.
"""
input_data: Final[str | ResponseInputParam | None] = data.get("input")
if input_data is None:
if not isinstance(input_data, (str, list)):
return data
structured_messages: Final = self.get_structured_messages(data)
raw_tools: Final = data.get("tools")
original_tools: Final[tuple[Mapping[str, object], ...]] = (
@ -173,94 +368,93 @@ class OpenAIResponsesHandler(BaseTranslation):
flattened_tool_groups: Final = tuple(
form.chat_tools for form in LiteLLMCompletionResponsesConfig.responses_tools_to_chat_forms(original_tools)
)
flattened_tools: Final = tuple(
cast(ChatCompletionToolParam, tool) # cast-ok: mcp tools ride along in the guardrail's tool list
for group in flattened_tool_groups
for tool in group
)
tools_to_check: Final[list[ChatCompletionToolParam]] = list( # mutable-ok: guardrail inputs want a list
copy.deepcopy(flattened_tools)
)
# Handle simple string input
if isinstance(input_data, str):
inputs = GenericGuardrailAPIInputs(texts=[input_data])
if tools_to_check:
inputs["tools"] = tools_to_check
if structured_messages:
inputs["structured_messages"] = structured_messages
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
)
guardrailed_texts = guardrailed_inputs.get("texts", [])
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data
self._apply_guardrailed_tools_to_data(
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
)
verbose_proxy_logger.debug("OpenAI Responses API: Processed string input")
return data
# Handle list input (ResponseInputParam)
if not isinstance(input_data, list):
extracted: Final = self._extract_guardrail_inputs(data, input_data, flattened_tool_groups)
if not extracted.inputs.get("texts"):
return data
if structured_messages:
extracted.inputs["structured_messages"] = structured_messages
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
inputs=extracted.inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
)
self._apply_guardrailed_tools_to_data(
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
)
written_back: Final = self._written_back_request_fields(data, structured_messages, guardrailed_inputs)
if written_back is not None:
data["input"] = list(written_back.input) # mutable-ok: JSON body
if written_back.instructions is None:
data.pop("instructions", None)
else:
data["instructions"] = written_back.instructions # rebind-ok: data is an out-param
elif isinstance(input_data, str):
guardrailed_texts: Final = guardrailed_inputs.get("texts") or ()
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param
else:
await self._apply_guardrail_responses_to_input(
messages=input_data,
responses=guardrailed_inputs.get("texts") or (),
task_mappings=extracted.task_mappings,
)
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", data.get("input"))
return data
def _extract_guardrail_inputs(
self,
data: Mapping[str, object],
input_data: "str | ResponseInputParam",
flattened_tool_groups: Sequence[Sequence[Mapping[str, object]]],
) -> _ExtractedInputs:
texts_to_check: Final[list[str]] = []
images_to_check: Final[list[str]] = []
task_mappings: Final[list[tuple[int, int | None]]] = []
# Step 1: Extract all text content, images, and tools
for msg_idx, message in enumerate(input_data):
self._extract_input_text_and_images(
message=message,
msg_idx=msg_idx,
texts_to_check=texts_to_check,
images_to_check=images_to_check,
task_mappings=task_mappings,
tools_to_check: Final[list[ChatCompletionToolParam]] = list( # mutable-ok: guardrail inputs want a list
copy.deepcopy(
tuple(
cast(ChatCompletionToolParam, tool) # cast-ok: mcp tools ride along in the guardrail's tool list
for group in flattened_tool_groups
for tool in group
)
)
)
if isinstance(input_data, str):
texts_to_check.append(input_data)
else:
for msg_idx, message in enumerate(input_data):
self._extract_input_text_and_images(
message=message,
msg_idx=msg_idx,
texts_to_check=texts_to_check,
images_to_check=images_to_check,
task_mappings=task_mappings,
)
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
if images_to_check:
inputs["images"] = images_to_check
if tools_to_check:
inputs["tools"] = tools_to_check
model: Final = data.get("model")
if isinstance(model, str):
inputs["model"] = model
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings))
# Step 2: Apply guardrail to all texts in batch
if texts_to_check:
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
if images_to_check:
inputs["images"] = images_to_check
if tools_to_check:
inputs["tools"] = tools_to_check
if structured_messages:
inputs["structured_messages"] = structured_messages
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
)
guardrailed_texts = guardrailed_inputs.get("texts", [])
self._apply_guardrailed_tools_to_data(
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
)
# Step 3: Map guardrail responses back to original input structure
await self._apply_guardrail_responses_to_input(
messages=input_data,
responses=guardrailed_texts,
task_mappings=task_mappings,
)
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", input_data)
return data
@staticmethod
def _written_back_request_fields(
data: Mapping[str, object],
structured_messages: Sequence[AllMessageValues] | None,
guardrailed_inputs: GenericGuardrailAPIInputs,
) -> _RequestFields | None:
guardrailed: Final = guardrailed_inputs.get("structured_messages")
if guardrailed is None or guardrailed is structured_messages:
return None
return _patch_or_convert_request_fields(
data.get("input"),
data.get("instructions"),
structured_messages or (),
guardrailed,
)
def extract_request_tool_names(self, data: dict) -> list[str]:
"""Extract tool names from Responses API request (tools[].name for function
@ -331,8 +525,8 @@ class OpenAIResponsesHandler(BaseTranslation):
async def _apply_guardrail_responses_to_input(
self,
messages: Any, # Can be List[Dict[str, Any]] or ResponseInputParam
responses: list[str],
task_mappings: list[tuple[int, int | None]],
responses: Sequence[str],
task_mappings: Sequence[tuple[int, int | None]],
) -> None:
"""
Apply guardrail responses back to input messages.

View file

@ -26,6 +26,7 @@ from copy import deepcopy
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, Union, cast, get_args
from urllib.parse import urlsplit
from litellm._logging import _redact_string
from litellm._uuid import uuid
@ -60,6 +61,7 @@ if TYPE_CHECKING:
from litellm.types.utils import TokenCountResponse
from litellm.constants import (
AZURE_OPENAI_AUDIO_PROVIDERS,
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT,
)
@ -984,6 +986,12 @@ def mock_completion(
_OPENAI_DEFAULT_API_BASE: Final = "https://api.openai.com/v1"
_OPENAI_API_HOST: Final = "api.openai.com"
def _is_openai_backed_api_base(api_base: str) -> bool:
hostname: Final = urlsplit(api_base).hostname
return hostname is not None and (hostname == _OPENAI_API_HOST or hostname.endswith(f".{_OPENAI_API_HOST}"))
def _resolve_openai_api_base(api_base: str | None) -> str:
@ -1053,7 +1061,7 @@ def responses_api_bridge_check(
# natively by Chat Completions with reasoning on, so custom-only requests stay on
# chat and keep their native custom tool_call response shape.
# - The UNSET-effort arm only fires against endpoints known to enforce that
# constraint (the default OpenAI endpoint, or Azure OpenAI where api_base is
# constraint (any api.openai.com host, or Azure OpenAI where api_base is
# always set): chat-only OpenAI-compatible backends registered under the openai
# provider with a custom api_base and gpt-5.4+ model names serve tools without
# reasoning fine and have no /responses route, so they keep pre-existing
@ -1068,14 +1076,15 @@ def responses_api_bridge_check(
reasoning_active = reasoning_effort.get("effort") != "none" or reasoning_effort.get("summary") is not None
else:
reasoning_active = reasoning_effort != "none"
# The reasoning+tools constraint is enforced only by the real OpenAI endpoint (and Azure OpenAI).
# Resolve the effective base arg>global>env>default exactly as the chat handler does, so a custom
# base set via litellm.api_base or OPENAI_BASE_URL/OPENAI_API_BASE isn't misread as the default and
# bridged to a /responses route it lacks. A whitespace-only base collapses to the default too.
resolved_api_base: Final = _resolve_openai_api_base(api_base)
on_constraint_enforcing_endpoint: Final = custom_llm_provider == "azure" or resolved_api_base.strip() in (
"",
_OPENAI_DEFAULT_API_BASE,
# The reasoning+tools constraint is enforced by the real OpenAI backend behind any api.openai.com
# host (the default URL or a PrivateLink hostname such as <region>.privatelink.api.openai.com) and
# by Azure OpenAI. Resolve the effective base arg>global>env>default exactly as the chat handler
# does, so a custom base set via litellm.api_base or OPENAI_BASE_URL/OPENAI_API_BASE isn't misread
# as the default and bridged to a /responses route it lacks. A whitespace-only base collapses to
# the default too.
resolved_api_base: Final = _resolve_openai_api_base(api_base).strip()
on_constraint_enforcing_endpoint: Final = (
custom_llm_provider == "azure" or resolved_api_base == "" or _is_openai_backed_api_base(resolved_api_base)
)
if (
custom_llm_provider in ("openai", "azure")
@ -7769,7 +7778,7 @@ def transcription(
provider=LlmProviders(custom_llm_provider),
)
if custom_llm_provider == "azure" and provider_config is None:
if custom_llm_provider in AZURE_OPENAI_AUDIO_PROVIDERS and provider_config is None:
# azure configs
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
@ -8056,7 +8065,10 @@ def speech(
custom_llm_provider=custom_llm_provider,
)
response: HttpxBinaryResponseContent | Coroutine[object, object, HttpxBinaryResponseContent] | None = None
if custom_llm_provider == "openai" or custom_llm_provider in litellm.openai_compatible_providers:
if custom_llm_provider == "openai" or (
custom_llm_provider in litellm.openai_compatible_providers
and custom_llm_provider not in AZURE_OPENAI_AUDIO_PROVIDERS
):
if voice is None or not (isinstance(voice, str)):
raise litellm.BadRequestError(
message="'voice' is required to be passed as a string for OpenAI TTS",
@ -8110,7 +8122,7 @@ def speech(
aspeech=aspeech,
shared_session=shared_session,
)
elif custom_llm_provider == "azure":
elif custom_llm_provider in AZURE_OPENAI_AUDIO_PROVIDERS:
# Check if this is Azure Speech Service (Cognitive Services TTS)
if model.startswith("speech/"):
from litellm.llms.azure.text_to_speech.transformation import (

View file

@ -29277,6 +29277,75 @@
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-6-astra": {
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
"cache_creation_input_token_cost_above_272k_tokens_flex": 1.25e-05,
"cache_creation_input_token_cost_above_272k_tokens_priority": 5e-05,
"cache_creation_input_token_cost_flex": 6.25e-06,
"cache_creation_input_token_cost_priority": 2.5e-05,
"cache_read_input_token_cost": 1e-06,
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
"cache_read_input_token_cost_above_272k_tokens_flex": 1e-06,
"cache_read_input_token_cost_above_272k_tokens_priority": 4e-06,
"cache_read_input_token_cost_flex": 5e-07,
"cache_read_input_token_cost_priority": 2e-06,
"input_cost_per_token": 1e-05,
"input_cost_per_token_above_272k_tokens": 2e-05,
"input_cost_per_token_above_272k_tokens_flex": 1e-05,
"input_cost_per_token_above_272k_tokens_priority": 4e-05,
"input_cost_per_token_batches": 5e-06,
"input_cost_per_token_flex": 5e-06,
"input_cost_per_token_priority": 2e-05,
"litellm_provider": "openai",
"max_input_tokens": 922000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5e-05,
"output_cost_per_token_above_272k_tokens": 7.5e-05,
"output_cost_per_token_above_272k_tokens_flex": 3.75e-05,
"output_cost_per_token_above_272k_tokens_priority": 0.00015,
"output_cost_per_token_batches": 2.5e-05,
"output_cost_per_token_flex": 2.5e-05,
"output_cost_per_token_priority": 0.0001,
"regional_processing_uplift_multiplier_eu": 1.1,
"regional_processing_uplift_multiplier_us": 1.1,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_cache_breakpoint": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_xhigh_reasoning_effort": true
},
"gpt-5.6": {
"cache_creation_input_token_cost": 5e-06,
"cache_creation_input_token_cost_above_272k_tokens": 1e-05,
@ -52911,6 +52980,11 @@
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
"output_cost_per_token": 3.3e-05,
"output_cost_per_token_above_272k_tokens": 4.95e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
"search_context_size_medium": 0.012
},
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@ -52933,7 +53007,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_web_search": true
},
"bedrock_mantle/openai.gpt-5.6-terra": {
"input_cost_per_token": 2.2e-06,
@ -52944,6 +53019,11 @@
"cache_read_input_token_cost_above_272k_tokens": 4.4e-07,
"output_cost_per_token": 1.32e-05,
"output_cost_per_token_above_272k_tokens": 1.98e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
"search_context_size_medium": 0.012
},
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@ -52966,7 +53046,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_web_search": true
},
"bedrock_mantle/openai.gpt-5.6-cyber": {
"input_cost_per_token": 1.375e-05,
@ -53005,6 +53086,11 @@
"cache_read_input_token_cost_above_272k_tokens": 4.4e-08,
"output_cost_per_token": 1.32e-06,
"output_cost_per_token_above_272k_tokens": 1.98e-06,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
"search_context_size_medium": 0.012
},
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@ -53027,7 +53113,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_web_search": true
},
"us.openai.gpt-5.6-sol": {
"input_cost_per_token": 4.4e-06,
@ -53192,6 +53279,11 @@
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
"output_cost_per_token": 3.3e-05,
"output_cost_per_token_above_272k_tokens": 4.95e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
"search_context_size_medium": 0.012
},
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@ -53213,7 +53305,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_web_search": true
},
"bedrock_mantle/openai.gpt-5.4": {
"input_cost_per_token": 2.75e-06,
@ -53222,6 +53315,11 @@
"cache_read_input_token_cost_above_272k_tokens": 5.5e-07,
"output_cost_per_token": 1.65e-05,
"output_cost_per_token_above_272k_tokens": 2.475e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
"search_context_size_medium": 0.012
},
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@ -53243,7 +53341,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_web_search": true
},
"bedrock_mantle/google.gemma-4-31b": {
"input_cost_per_token": 1.4e-07,

View file

@ -3,7 +3,7 @@ from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, cast
from typing import TYPE_CHECKING, Final, Literal, cast
from fastapi import HTTPException
from starlette.datastructures import Headers
@ -305,6 +305,12 @@ def _admission_failure_fallback(
raise exc
@dataclass(frozen=True, slots=True)
class MCPServerAccess:
server_ids: tuple[str, ...]
scope: Literal["unscoped", "scoped", "unresolved"] = "unscoped"
@dataclass(frozen=True, slots=True)
class DcrBridgeTarget:
"""The single DCR-bridge server a request targets, paired with the exact name the caller
@ -1456,6 +1462,18 @@ class MCPRequestHandler:
*,
keyless_source: bool = False,
) -> list[str]:
access: Final = await MCPRequestHandler.get_mcp_server_access(
user_api_key_auth,
keyless_source=keyless_source,
)
return list(access.server_ids)
@staticmethod
async def get_mcp_server_access(
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
keyless_source: bool = False,
) -> MCPServerAccess:
"""
Get list of allowed MCP servers for the given user/key based on permissions.
@ -1478,13 +1496,17 @@ class MCPRequestHandler:
"""
from litellm.proxy.proxy_server import general_settings
key_object_permission: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
try:
# A keyless admitted subject resolves per source BEFORE any single-source rule here. Ordering
# matters: the no_mcp_servers opt-out below reads the caller's own object_permission, so above
# this branch a user's own opt-out would wrongly zero their TEAMS' grants too (each source is
# independent; an opt-out silences only its own source, inside the recursive call).
if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
return await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)
return MCPServerAccess(
server_ids=tuple(await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)),
)
# Get allowed servers from key and team
allowed_mcp_servers_for_key = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
@ -1492,7 +1514,7 @@ class MCPRequestHandler:
# The key explicitly opted out of every MCP server. This overrides
# team inheritance and additive grants (mirrors no-default-models).
if SpecialMCPServerNames.no_mcp_servers.value in allowed_mcp_servers_for_key:
return []
return MCPServerAccess(server_ids=(), scope="scoped")
allowed_mcp_servers_for_team = await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_api_key_auth)
@ -1572,7 +1594,7 @@ class MCPRequestHandler:
"require_end_user_mcp_access_defined=True and end_user %s has no MCP permissions - blocking MCP access",
user_api_key_auth.end_user_id,
)
return []
return MCPServerAccess(server_ids=(), scope="scoped")
#########################################################
# Check agent permissions if agent_id is set on the key
@ -1601,14 +1623,22 @@ class MCPRequestHandler:
#########################################################
# Apply org-level ceiling if org_id is set
#########################################################
allowed_mcp_servers = await MCPRequestHandler._apply_primary_org_ceiling(
allowed_mcp_servers, org_restricts = await MCPRequestHandler._apply_primary_org_ceiling(
allowed_mcp_servers,
user_api_key_auth,
has_lower_level_mcp_restrictions,
keyless_source=keyless_source,
)
return list(set(allowed_mcp_servers))
declares_key_mcp_scope: Final = getattr(key_object_permission, "mcp_servers", None) is not None
return MCPServerAccess(
server_ids=tuple(set(allowed_mcp_servers)),
scope=(
"scoped"
if has_lower_level_mcp_restrictions or org_restricts or declares_key_mcp_scope
else "unscoped"
),
)
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
# A ceiling we KNOW exists and cannot read. Denying is the only answer that does not
@ -1616,7 +1646,10 @@ class MCPRequestHandler:
verbose_logger.warning("Denying MCP access, entitlement unreadable: %s", e)
else:
verbose_logger.warning("Failed to get allowed MCP servers: %s", e)
return []
return MCPServerAccess(
server_ids=(),
scope="scoped" if getattr(key_object_permission, "mcp_servers", None) is not None else "unresolved",
)
@staticmethod
async def _apply_primary_org_ceiling(
@ -1624,7 +1657,7 @@ class MCPRequestHandler:
user_api_key_auth: UserAPIKeyAuth | None,
has_lower_level_mcp_restrictions: bool,
keyless_source: bool = False,
) -> list[str]:
) -> tuple[list[str], bool]:
"""Cap the resolved server list by this caller's org ceiling: an explicit org list intersects
lower-level restrictions (else becomes the ceiling); no org or an empty list leaves it unchanged.
@ -1638,7 +1671,7 @@ class MCPRequestHandler:
cannot be read raises out of ``_get_allowed_mcp_servers_for_org`` and never arrives here as
``None``, so key auth cannot silently shed a ceiling an operator did configure."""
if not (user_api_key_auth and user_api_key_auth.org_id):
return allowed_mcp_servers
return allowed_mcp_servers, False
allowed_mcp_servers_for_org: Final = await MCPRequestHandler._get_allowed_mcp_servers_for_org(user_api_key_auth)
if allowed_mcp_servers_for_org is None:
verbose_logger.warning(
@ -1646,9 +1679,9 @@ class MCPRequestHandler:
user_api_key_auth.org_id,
"denying (keyless admitted subject)" if keyless_source else "leaving uncapped (key auth)",
)
return [] if keyless_source else allowed_mcp_servers
return ([] if keyless_source else allowed_mcp_servers), False
if len(allowed_mcp_servers_for_org) == 0:
return allowed_mcp_servers
return allowed_mcp_servers, False
if has_lower_level_mcp_restrictions or keyless_source:
# Org can only cap lower-level restrictions. A keyless admitted source ALWAYS takes this
# arm: its model unions GRANTS, so an org list may only narrow a source, never become one.
@ -1657,7 +1690,7 @@ class MCPRequestHandler:
# No lower-level restrictions → org list becomes the ceiling.
capped = allowed_mcp_servers_for_org
verbose_logger.debug("Applied org ceiling filter. Final allowed servers: %s", capped)
return capped
return capped, True
@staticmethod
def _scoped_source_auth(

View file

@ -55,6 +55,7 @@ from litellm.litellm_core_utils.url_utils import SSRFError, async_safe_get
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
MCPServerAccess,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
@ -2958,7 +2959,13 @@ class MCPServerManager:
return None
return user_api_key_auth.mcp_session_resource_server_id
async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth | None = None) -> list[str]:
async def get_allowed_mcp_servers(
self,
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
access: MCPServerAccess | None = None,
general_settings: Mapping[str, object] | None = None,
) -> list[str]:
"""
Get the allowed MCP Servers for the user.
@ -2967,6 +2974,9 @@ class MCPServerManager:
2. If admin and no object_permission, return all servers
3. Otherwise, use standard permission checks
"""
from litellm.proxy.proxy_server import general_settings as proxy_general_settings
resolved_general_settings: Final = proxy_general_settings if general_settings is None else general_settings
allow_all_server_ids: Final = self.get_allow_all_keys_server_ids()
# A keyless admitted subject is resolved per grant source, and channel decisions that are
@ -3007,11 +3017,16 @@ class MCPServerManager:
# whole registry, for keys AND admitted session subjects alike (one predicate owns the
# question). Seeded into the union rather than returned early so the session resource
# scope below still bounds a per-server envelope held by an admin.
combined_servers: Final = (
set(self.get_registry().keys())
if await MCPRequestHandler.admin_view_unscoped(user_api_key_auth)
else set(await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth))
admin_unscoped: Final = await MCPRequestHandler.admin_view_unscoped(user_api_key_auth)
resolved_access: Final = (
MCPServerAccess(server_ids=())
if admin_unscoped
else access or await MCPRequestHandler.get_mcp_server_access(user_api_key_auth)
)
resolved_server_ids: Final = (
set(self.get_registry().keys()) if admin_unscoped else set(resolved_access.server_ids)
)
combined_servers: Final = set(resolved_server_ids)
verbose_logger.debug("Allowed MCP Servers for user api key auth: %s", combined_servers)
combined_servers.update(
await self.operator_open_server_ids(
@ -3052,6 +3067,18 @@ class MCPServerManager:
]
combined_servers.update(delegate_server_ids)
restrict_allow_all: Final = (
resolved_general_settings.get("mcp_allow_all_keys_respects_mcp_scope", False)
and user_api_key_auth is not None
and user_api_key_auth.via_virtual_key
and resolved_access.scope != "unscoped"
)
if restrict_allow_all:
combined_servers.difference_update(
set(allow_all_server_ids)
- resolved_server_ids
- (set(submitted_server_ids) if resolved_access.scope != "unresolved" else set())
)
if len(combined_servers) == 0:
verbose_logger.debug("No allowed MCP Servers found for user api key auth.")
scope = MCPServerManager._admitted_session_resource_scope(user_api_key_auth)

View file

@ -1216,7 +1216,7 @@ class GenerateKeyRequest(KeyRequestBase):
organization_id: str | None = None
project_id: str | None = None
@field_validator("team_id", "organization_id", mode="before")
@field_validator("team_id", "organization_id", "project_id", mode="before")
@classmethod
def treat_cleared_id_as_unset(cls, v: object) -> object:
if v == "":
@ -1930,6 +1930,13 @@ class NewTeamRequest(TeamBase):
model_config = ConfigDict(protected_namespaces=())
@field_validator("team_id", mode="before")
@classmethod
def treat_blank_team_id_as_unset(cls, v: object) -> object:
if isinstance(v, str) and not v.strip():
return None
return v
class GlobalEndUsersSpend(LiteLLMPydanticObjectBase):
api_key: str | None = None
@ -2601,9 +2608,9 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata.",
)
missing_session_id: Literal["generate", "reject"] | None = Field(
missing_session_id: Literal["generate", "reject", "omit"] | None = Field(
None,
description="What to do with LLM API requests that carry no session id (x-litellm-session-id header, metadata.session_id, etc.). 'generate' stamps one id into litellm_session_id, litellm_trace_id and metadata.session_id so SpendLogs and logging callbacks agree; 'reject' returns 400. Unset keeps the legacy behavior where SpendLogs falls back to the trace id while callbacks get no session id.",
description="What to do with LLM API requests that carry no session id (x-litellm-session-id header, metadata.session_id, etc.). 'generate' stamps one id into litellm_session_id, litellm_trace_id and metadata.session_id so SpendLogs and logging callbacks agree; 'reject' returns 400; 'omit' leaves SpendLogs.session_id null, matching callbacks such as Langfuse that only record a client-established metadata.session_id. Unset keeps the legacy behavior where SpendLogs falls back to the trace id while callbacks get no session id.",
)
enable_public_model_hub: bool = Field(
default=False,
@ -4232,6 +4239,8 @@ class TeamAccessGroupModelGrant(LiteLLMPydanticObjectBase):
access_group_id: str
access_group_name: str
models: tuple[str, ...]
mcp_server_ids: tuple[str, ...] = ()
agent_ids: tuple[str, ...] = ()
class TeamInfoResponseObjectTeamTable(LiteLLM_TeamTable):

View file

@ -939,35 +939,32 @@ async def make_agent_public(
if agent is None:
raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found")
if litellm.public_agent_groups is None:
litellm.public_agent_groups = []
# handle duplicates
if not AGENT_REGISTRY.ids_for_agent(agent.agent_id).isdisjoint(litellm.public_agent_groups):
config: Final = await proxy_config.get_config()
current_public_agent_groups: Final = list(litellm.public_agent_groups or [])
if not AGENT_REGISTRY.ids_for_agent(agent.agent_id).isdisjoint(current_public_agent_groups):
raise HTTPException(
status_code=400,
detail=f"Agent with name {agent.agent_name} already in public agent groups",
)
litellm.public_agent_groups.append(agent.agent_id)
updated_public_agent_groups: Final = [*current_public_agent_groups, agent.agent_id]
# Load existing config
config: Final = await proxy_config.get_config()
# Update config with new settings
if "litellm_settings" not in config or config["litellm_settings"] is None:
config["litellm_settings"] = {}
config["litellm_settings"]["public_agent_groups"] = litellm.public_agent_groups
config["litellm_settings"]["public_agent_groups"] = updated_public_agent_groups
# Save the updated config
await proxy_config.save_config(new_config=config)
litellm.public_agent_groups = updated_public_agent_groups
verbose_proxy_logger.debug(
"Updated public agent groups to: %s by user: %s", litellm.public_agent_groups, user_api_key_dict.user_id
"Updated public agent groups to: %s by user: %s", updated_public_agent_groups, user_api_key_dict.user_id
)
return {
"message": "Successfully updated public agent groups",
"public_agent_groups": litellm.public_agent_groups,
"public_agent_groups": updated_public_agent_groups,
"updated_by": user_api_key_dict.user_id,
}
except HTTPException:

View file

@ -16,6 +16,10 @@ if TYPE_CHECKING:
from litellm.proxy._types import EnterpriseLicenseData
AUTO_ROUTER_LICENSE_FEATURE: Final = "auto_router"
HEURISTIC_V2_LICENSE_REMEDY: Final = "A LiteLLM license with the 'auto_router' feature lifts the limit."
class LicenseCheck:
"""
- Check if license in env
@ -149,6 +153,19 @@ class LicenseCheck:
return False
return team_count > _max_teams_in_license
def heuristic_v2_router_limit(self) -> int | None:
"""
How many heuristic_v2 auto-routers this proxy may hold: unlimited (None) only when the
signed license lists the auto_router feature, otherwise one. A license verified through
the API carries no feature list, so it does not lift the limit either.
"""
if self.airgapped_license_data is None:
return 1
allowed_features: Final = self.airgapped_license_data.get("allowed_features")
if isinstance(allowed_features, list) and AUTO_ROUTER_LICENSE_FEATURE in allowed_features:
return None
return 1
def verify_license_without_api_request(self, public_key, license_key):
try:
from cryptography.hazmat.primitives import hashes
@ -179,19 +196,21 @@ class LicenseCheck:
# Decode and parse the data
license_data: Final = json.loads(message.decode())
self.airgapped_license_data = EnterpriseLicenseData(**license_data)
# debug information provided in license data
verbose_proxy_logger.debug("License data: %s", license_data)
# Check expiration date
expiration_date: Final = datetime.strptime(license_data["expiration_date"], "%Y-%m-%d")
if expiration_date < datetime.now():
self.airgapped_license_data = None
return False, "License has expired"
self.airgapped_license_data = EnterpriseLicenseData(**license_data)
return True
except Exception as e:
self.airgapped_license_data = None
verbose_proxy_logger.debug(
"litellm.proxy.auth.litellm_license.py::verify_license_without_api_request - Unable to verify License locally. - %s",
e,

View file

@ -124,6 +124,42 @@ def add_missing_query_params(url: str, params: Mapping[str, str | int | float])
return urllib.parse.urlunsplit(parsed._replace(query=query))
LIBPQ_VERIFY_SSLMODES: Final[frozenset[str]] = frozenset({"verify-ca", "verify-full"})
def translate_libpq_ssl_params(url: str) -> str:
"""Rewrite libpq's certificate-verification params into Prisma's dialect.
Prisma's engine only knows ``sslmode=disable|prefer|require``, ``sslcert``
(the CA bundle) and ``sslaccept=strict``. It silently discards
``sslrootcert`` and downgrades ``sslmode=verify-ca`` / ``verify-full`` to
``prefer``, so a URL copied from libpq / RDS docs connects over TLS with no
certificate check at all. ``verify-ca`` and ``verify-full`` both become
``require`` (Prisma has no CA-only mode), ``sslrootcert`` becomes
``sslcert``, and either one turns on ``sslaccept=strict`` (chain and
hostname), matching libpq where a root cert makes ``require`` verify.
Prisma params the operator pinned themselves win; anything else is left
untouched.
"""
parsed: Final = urllib.parse.urlsplit(url)
pairs: Final = tuple(urllib.parse.parse_qsl(parsed.query, keep_blank_values=True))
keys: Final = frozenset(key for key, _ in pairs)
wants_verify: Final = any(key == "sslmode" and value in LIBPQ_VERIFY_SSLMODES for key, value in pairs)
if not wants_verify and "sslrootcert" not in keys:
return url
translated: Final = tuple(
("sslmode", "require") if key == "sslmode" and value in LIBPQ_VERIFY_SSLMODES else (key, value)
for key, value in pairs
if key != "sslrootcert"
)
root_cert: Final = tuple(
("sslcert", value) for key, value in pairs if key == "sslrootcert" and "sslcert" not in keys
)
strict: Final = () if "sslaccept" in keys else (("sslaccept", "strict"),)
query: Final = urllib.parse.urlencode(translated + root_cert + strict)
return urllib.parse.urlunsplit(parsed._replace(query=query))
def reader_shareable_params(params: Mapping[str, str | int | float]) -> Mapping[str, str | int | float]:
"""Return the subset of ``params`` the read replica is allowed to inherit."""
return MappingProxyType({key: value for key, value in params.items() if key in CONNECTION_PARAM_KEYS})
@ -403,6 +439,11 @@ class DatabaseURLSettings(BaseSettings):
self._raise_for_unsupported_scheme()
wrote_writer: Final = self.apply_writer_url_to_env()
for env_var in ("DATABASE_URL", "DIRECT_URL"):
url = os.environ.get(env_var)
if url:
os.environ[env_var] = translate_libpq_ssl_params(url)
# DATABASE_DISABLE_PREPARED_STATEMENTS maps to Prisma's `pgbouncer=true`
# URL param, same as the CLI's `database_disable_prepared_statements`
# config key. An explicit `pgbouncer` value already on the URL wins.
@ -418,7 +459,7 @@ class DatabaseURLSettings(BaseSettings):
reader_url: Final = self.build_reader_url() or self.database_url_read_replica
if reader_url is not None:
os.environ["DATABASE_URL_READ_REPLICA"] = add_missing_query_params(
reader_url,
translate_libpq_ssl_params(reader_url),
connection_params_from_url(os.environ.get("DATABASE_URL", "")),
)

View file

@ -916,9 +916,10 @@ class CompresrGuardrail(CustomGuardrail):
def _mirror_texts_channel(input_texts: object, applied: _CompressionResult) -> list[object] | None:
"""Compressed content mirrored into the Responses `texts` channel.
The chat/Anthropic handlers round-trip ``structured_messages``; the
Responses translation cannot rebuild its input from chat messages and
instead writes back through ``texts``. This matches by value, so a
The chat/Anthropic/Responses handlers round-trip
``structured_messages``; translations without that round-trip write
back through ``texts``, so the compressed content is mirrored there
too. This matches by value, so a
replacement is applied only when it is unambiguous: one compression per
text, and every occurrence in ``texts`` accounted for by a compressed
target. Anything else is left uncompressed rather than risk a wrong or

View file

@ -50,6 +50,9 @@ if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
BYPASS_HEADER: Final = "x-headroom-bypass"
_STREAM_CONVERTIBLE_CALL_TYPES: Final = frozenset(
(CallTypes.completion, CallTypes.acompletion, CallTypes.responses, CallTypes.aresponses)
)
HEADROOM_RETRIEVE_TOOL_NAME: Final = "headroom_retrieve"
_HASH_PATTERN: Final = re.compile(r"hash=([a-f0-9]{24})")
_HASH_CACHE_TTL_SECONDS: Final = 15 * 60
@ -725,6 +728,10 @@ class HeadroomGuardrail(CustomGuardrail):
verbose_proxy_logger.debug("Headroom: %s header set; skipping compression", BYPASS_HEADER)
return inputs
if request_data.get("background"):
verbose_proxy_logger.debug("Headroom: background request; skipping compression")
return inputs
structured_messages: Final = inputs.get("structured_messages")
if not _is_object_list(structured_messages) or not structured_messages:
return inputs
@ -826,9 +833,9 @@ class HeadroomGuardrail(CustomGuardrail):
) -> dict[str, Any] | None: # mutable-ok: overrides CustomLogger hook whose contract is a plain dict
base_result: Final = await super().async_pre_call_deployment_hook(kwargs, call_type)
effective: Final = base_result if base_result is not None else kwargs
if call_type not in (CallTypes.completion, CallTypes.acompletion):
if call_type not in _STREAM_CONVERTIBLE_CALL_TYPES:
return base_result
if not effective.get("stream"):
if not effective.get("stream") or effective.get("background"):
return base_result
if not has_headroom_retrieve_tool(effective.get("tools")):
return base_result

View file

@ -168,11 +168,8 @@ class _ProxyDBLogger(CustomLogger):
"custom_llm_provider"
) or request_data.get("custom_llm_provider", "")
# Propagate standard_logging_object and litellm_trace_id from the
# Logging instance so that _get_session_id_for_spend_log uses the same
# trace_id that Langfuse received (via async_failure_handler).
# Without this, the DB session_id would be a random UUID that doesn't
# match the Langfuse trace_id, making failed requests unsearchable.
# Propagate standard_logging_object and litellm_trace_id from the Logging
# instance so the failure row carries the same trace_id Langfuse received.
_litellm_logging_obj: Final = request_data.get("litellm_logging_obj")
if _litellm_logging_obj is not None:
if not request_data.get("standard_logging_object"):

View file

@ -25,6 +25,7 @@ from litellm.constants import (
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
SESSION_ID_GENERATED_METADATA_KEY,
SESSION_ID_OMITTED_METADATA_KEY,
)
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
@ -733,12 +734,18 @@ def apply_missing_session_id_policy(
general_settings: Mapping[str, object] | None,
request: Request,
) -> None:
for metadata_key in ("metadata", "litellm_metadata"):
if isinstance(client_metadata := data.get(metadata_key), dict):
client_metadata.pop(SESSION_ID_OMITTED_METADATA_KEY, None)
metadata: Final = data.get(_metadata_variable_name)
policy: Final = general_settings.get("missing_session_id") if general_settings else None
if policy is None or not _is_llm_inference_route(request):
return
metadata: Final = data.get(_metadata_variable_name)
if not isinstance(metadata, dict):
return
if policy == "omit":
metadata[SESSION_ID_OMITTED_METADATA_KEY] = True
return
if data.get("litellm_session_id") or metadata.get("session_id"):
return
match policy:
@ -760,7 +767,8 @@ def apply_missing_session_id_policy(
)
case _:
verbose_proxy_logger.warning(
"Ignoring unknown general_settings.missing_session_id=%r; expected 'generate' or 'reject'", policy
"Ignoring unknown general_settings.missing_session_id=%r; expected 'generate', 'reject' or 'omit'",
policy,
)

View file

@ -426,6 +426,11 @@ async def add_new_user_to_default_team(
await asyncio.gather(*tasks, return_exceptions=True)
async def _fetch_user_team_ids(user_id: str, prisma_client: "PrismaClient") -> tuple[str, ...]:
user_row: Final = await _user_table(prisma_client).find_unique(where={"user_id": user_id})
return tuple(user_row.teams) if user_row is not None else ()
@router.post(
"/user/new",
tags=["Internal User management"],
@ -580,6 +585,11 @@ async def new_user(
)
user_id: Final = cast(str | None, response.get("user_id", None))
attached_team_ids: Final = (
await _fetch_user_team_ids(user_id=user_id, prisma_client=prisma_client)
if user_id is not None and (_team_id is not None or teams is not None)
else None
)
if organization_ids is not None and user_id is not None:
await _add_user_to_organizations(
@ -596,6 +606,8 @@ async def new_user(
response_dict[key] = value
response_dict["key"] = response.get("token", "")
if attached_team_ids is not None:
response_dict["teams"] = list(attached_team_ids)
new_user_response: Final = NewUserResponse.model_validate(response_dict)

View file

@ -13,10 +13,11 @@ model/{model_id}/update - PATCH endpoint for model update.
import asyncio
import datetime
import json
from collections.abc import Awaitable, Mapping, Sequence
from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from json import JSONDecodeError
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeVar, cast
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator
@ -50,6 +51,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_utils import reject_server_owned_wif_params
from litellm.proxy.auth.litellm_license import HEURISTIC_V2_LICENSE_REMEDY
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.config_sync_pubsub import (
coordination_redis_cache,
@ -102,6 +104,9 @@ from litellm.router_strategy.complexity_router import (
from litellm.router_utils.auto_router_model_naming import (
STRATEGY_ROUTER_PARAM_FIELDS,
carries_complexity_router_settings,
count_heuristic_v2_routers,
heuristic_v2_limit_violation,
uses_heuristic_v2_classifier,
validate_complexity_router_config_placement,
validate_complexity_router_config_write,
validate_strategy_router_model_write,
@ -162,6 +167,8 @@ class _ProxyModelTable(Protocol):
def find_many(self, *, where: Mapping[str, object]) -> Awaitable[Sequence[_ProxyModelRow]]: ...
def create(self, *, data: Mapping[str, object]) -> Awaitable[_ProxyModelRow]: ...
def update(
self, *, where: Mapping[str, object], data: Mapping[str, object]
) -> Awaitable[_ProxyModelRow | None]: ...
@ -175,6 +182,9 @@ class _TxModelTables(Protocol):
litellm_proxymodeltable: _ProxyModelTable
_RowT = TypeVar("_RowT")
class _ExistingModelRow(Protocol):
@property
def litellm_params(self) -> Mapping[str, object]: ...
@ -295,6 +305,66 @@ def _reject_non_admin_blocked_flag_on_create(
)
HEURISTIC_V2_SLOT_LOCK_KEY: Final = 5_872_301
_HEURISTIC_V2_LOCK_SQL: Final = "SELECT 1 AS locked FROM pg_advisory_xact_lock($1)"
_HEURISTIC_V2_DB_ROWS_SQL: Final = """
SELECT count(*)::int AS held FROM "LiteLLM_ProxyModelTable"
WHERE model_id <> $1
AND (CASE jsonb_typeof(litellm_params) WHEN 'string' THEN (litellm_params #>> '{}')::jsonb ELSE litellm_params END)
-> 'complexity_router_config' ->> 'classifier_type' = 'heuristic_v2'
"""
def _effective_complexity_router_config(
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
) -> object:
"""The complexity config a write leaves on the row: the incoming one when the write carries it, else the stored one."""
incoming: Final = None if incoming_params is None else incoming_params.complexity_router_config
if incoming is not None or existing_params is None:
return incoming
return existing_params.complexity_router_config
@asynccontextmanager
async def _heuristic_v2_slot(
prisma_client: PrismaClient, *, effective_config: object, model_id: str | None
) -> AsyncGenerator[_ProxyModelTable, None]:
"""Hand out the model table to write through while the row's claim on a heuristic_v2 slot is settled.
A write that leaves the row on classifier_type heuristic_v2 under a limited license runs
inside one transaction that takes an advisory lock in its own statement before counting
(a statement's snapshot predates anything it locks), so pods cannot both pass the count:
the DB rows (any pod, either JSON shape) plus this proxy's config.yaml routers are judged
against the license limit and the write is refused with a 403 before it happens. The row
being edited keeps its own slot through ``model_id``. Every other write, and every write on
an unlimited license, goes through the repository table with no lock. Only the row write
itself may run inside: anything that needs a second connection (the team model bookkeeping)
must wait until the transaction has committed and the lock is released. The transaction
writes bypass the repository's publish-on-write, so the config change is published once
after commit, the way delete_team_models does.
"""
from litellm.proxy.proxy_server import _license_check, llm_router
limit: Final = _license_check.heuristic_v2_router_limit()
if limit is None or not uses_heuristic_v2_classifier(effective_config):
yield _proxy_model_table(prisma_client)
return
async with prisma_client.db.tx() as tx_ctx:
tables: Final[_TxModelTables] = tx_ctx
await tx_ctx.query_raw(_HEURISTIC_V2_LOCK_SQL, HEURISTIC_V2_SLOT_LOCK_KEY)
rows: Sequence[Mapping[str, object]] = await tx_ctx.query_raw(_HEURISTIC_V2_DB_ROWS_SQL, model_id or "")
db_held: Final = rows[0].get("held") if rows else 0
config_rows: Final = () if llm_router is None else tuple(llm_router.config_deployments())
held: Final = (db_held if isinstance(db_held, int) else 0) + count_heuristic_v2_routers(config_rows)
violation: Final = heuristic_v2_limit_violation(held=held + 1, limit=limit)
if violation is not None:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN, detail=f"{violation} {HEURISTIC_V2_LICENSE_REMEDY}"
)
yield tables.litellm_proxymodeltable
await publish_config_change(redis_cache=coordination_redis_cache(), object_type="litellm_proxymodeltable")
ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING: Final = "enforce_rpm_tpm_on_model_add"
_REQUIRED_RATE_LIMIT_FIELDS: Final = ("rpm", "tpm")
@ -747,22 +817,29 @@ async def patch_model(
)
requested_model_name: Final = patch_data.model_name
stored_model_name: str | None = None
async def write_row(update_data: PrismaCompatibleUpdateDBModel) -> _ProxyModelRow | None:
nonlocal stored_model_name
stored_model_name = update_data.get("model_name")
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
update_data["updated_at"] = cast(str, get_utc_datetime())
async with _heuristic_v2_slot(
prisma_client,
effective_config=_effective_complexity_router_config(
patch_data.litellm_params, db_model.litellm_params
),
model_id=model_id,
) as table:
return await table.update(where={"model_id": model_id}, data=update_data)
# Handle team model updates with proper alias management
update_data: Final = await _update_team_model_in_db(
updated_model: Final = await _update_team_model_in_db(
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
# Add metadata about update
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
update_data["updated_at"] = cast(str, get_utc_datetime())
# Perform partial update
updated_model: Final = await _proxy_model_table(prisma_client).update(
where={"model_id": model_id},
data=update_data,
write_row=write_row,
)
if updated_model is None:
@ -773,7 +850,6 @@ async def patch_model(
param=None,
)
stored_model_name: Final = update_data.get("model_name")
if (
stored_model_name is not None
and stored_model_name == requested_model_name
@ -1007,7 +1083,8 @@ async def _add_model_to_db(
prisma_client: PrismaClient,
new_encryption_key: str | None = None,
should_create_model_in_db: bool = True,
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
) -> "_ProxyModelRow | LiteLLM_ProxyModelTable":
# encrypt litellm params #
_litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True)
_original_litellm_model_name: Final = model_params.litellm_params.model
@ -1027,18 +1104,20 @@ async def _add_model_to_db(
if model_params.blocked is not None:
_data["blocked"] = model_params.blocked
_create_data: Final = cast("Mapping[str, object]", _data) # cast-ok: str-keyed json payload built just above
if should_create_model_in_db:
model_response = await ModelRepository(prisma_client).table.create(data=_create_data)
else:
model_response = LiteLLM_ProxyModelTable(**_data)
return model_response
if not should_create_model_in_db:
return LiteLLM_ProxyModelTable(**_data)
if slot is None:
return await _proxy_model_table(prisma_client).create(data=_create_data)
async with slot as table:
return await table.create(data=_create_data)
async def _add_team_model_to_db(
model_params: Deployment,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None":
slot: AbstractAsyncContextManager[_ProxyModelTable] | None = None,
) -> "_ProxyModelRow | LiteLLM_ProxyModelTable":
"""
If 'team_id' is provided,
@ -1069,6 +1148,7 @@ async def _add_team_model_to_db(
model_params=model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
slot=slot,
)
if original_model_name:
@ -1089,7 +1169,8 @@ async def _update_team_model_in_db(
patch_data: updateDeployment,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
) -> PrismaCompatibleUpdateDBModel:
write_row: Callable[[PrismaCompatibleUpdateDBModel], Awaitable[_RowT]],
) -> _RowT:
"""
Handle team model updates with proper alias management.
@ -1097,6 +1178,9 @@ async def _update_team_model_in_db(
- Creates unique internal model_name and team alias
- Adds model to team object
- Preserves team_public_model_name for external reference
The row is written through ``write_row`` before the team's model list is touched, so a
refused or failed write leaves the team as it was (the create path orders itself the same way).
"""
# Validate team_id if present in patch_data
from litellm.proxy.proxy_server import premium_user
@ -1108,9 +1192,7 @@ async def _update_team_model_in_db(
premium_user=premium_user,
)
# Validated before any write, beside the premium check the create path already runs
# here. The team ACL is updated below and autocommits, so a validator that raises
# further down would leave the team mutated and the deployment row never written.
# Validated before the row write, beside the premium check the create path already runs here.
#
# The merged view is what gets stored, so that is what has to satisfy the invariants.
# Validating the patch alone rejected a partial edit of an already valid deployment:
@ -1130,7 +1212,7 @@ async def _update_team_model_in_db(
# No team_id in patch, proceed with standard update
if patch_team_id is None:
return update_db_model(db_model=db_model, updated_patch=patch_data)
return await write_row(update_db_model(db_model=db_model, updated_patch=patch_data))
# Determine public model name
public_model_name: Final = _get_public_model_name(
@ -1149,11 +1231,14 @@ async def _update_team_model_in_db(
db_team_id: Final = db_model.model_info.team_id if db_model.model_info else None
is_new_team_assignment: Final = db_team_id != patch_team_id
# Team rows keep their internal UUID-based model_name; the public name lives in model_info
patch_data.model_name = f"model_name_{patch_team_id}_{uuid.uuid4()}" if is_new_team_assignment else None
row: Final = await write_row(update_db_model(db_model=db_model, updated_patch=patch_data))
if is_new_team_assignment:
await _setup_new_team_model_assignment(
team_id=patch_team_id,
public_model_name=public_model_name,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
)
else:
@ -1161,12 +1246,11 @@ async def _update_team_model_in_db(
team_id=patch_team_id,
public_model_name=public_model_name,
db_model=db_model,
patch_data=patch_data,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
return update_db_model(db_model=db_model, updated_patch=patch_data)
return row
def _get_public_model_name(
@ -1218,13 +1302,9 @@ def _get_public_model_name(
async def _setup_new_team_model_assignment(
team_id: str,
public_model_name: str,
patch_data: updateDeployment,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""Set up a new team model with unique name and team membership."""
unique_model_name: Final = f"model_name_{team_id}_{uuid.uuid4()}"
patch_data.model_name = unique_model_name
"""Register a newly team-assigned model's public name on the team."""
await team_model_add(
data=TeamModelAddRequest(
team_id=team_id,
@ -1414,7 +1494,6 @@ async def _update_existing_team_model_assignment(
team_id: str,
public_model_name: str,
db_model: Deployment,
patch_data: updateDeployment,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient | None,
) -> None:
@ -1438,9 +1517,6 @@ async def _update_existing_team_model_assignment(
old_public_name: Final = db_model.model_info.team_public_model_name if db_model.model_info else None
if old_public_name and public_model_name != old_public_name:
# Clear user-supplied public name from patch before any early return so the
# caller does not overwrite the internal UUID-based model_name in the DB.
patch_data.model_name = None
if prisma_client is None:
verbose_proxy_logger.warning(
"prisma_client not initialized; skipping public name update entirely to avoid orphaned entries"
@ -1488,10 +1564,6 @@ async def _update_existing_team_model_assignment(
# else: old_public_name == public_model_name (no rename needed)
# No team_model_add/delete calls required; public name is already registered
# Always clear patch_data.model_name to prevent caller from overwriting
# the internal UUID-based model_name in the DB with the user-supplied public name
patch_data.model_name = None
class ModelManagementAuthChecks:
"""
@ -2072,18 +2144,19 @@ async def add_new_model(
reload_outcome: ReconcileOutcome = ReconcileOutcome(still_desired=None, live_after=None)
try:
_original_litellm_model_name: Final = model_params.model_name
if model_params.model_info.team_id is None:
model_response = await _add_model_to_db(
model_params=priced_model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
else:
model_response = await _add_team_model_to_db(
model_params=priced_model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
add_model: Final = (
_add_model_to_db if model_params.model_info.team_id is None else _add_team_model_to_db
)
model_response = await add_model(
model_params=priced_model_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
slot=_heuristic_v2_slot(
prisma_client,
effective_config=priced_model_params.litellm_params.complexity_router_config,
model_id=priced_model_params.model_info.id,
),
)
reload_outcome = await proxy_config.add_deployment(
prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj
)
@ -2097,6 +2170,8 @@ async def add_new_model(
passed_model_info=priced_model_params.model_info,
)
except Exception as e:
if isinstance(e, HTTPException):
raise
verbose_proxy_logger.exception("Exception in add_new_model: %s", e)
else:
@ -2265,10 +2340,17 @@ async def update_model(
"updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
**({} if renamed_to is None else {"model_name": renamed_to}),
}
model_response: Final = await _proxy_model_table(prisma_client).update(
where={"model_id": _model_id},
data=_data,
)
async with _heuristic_v2_slot(
prisma_client,
effective_config=_effective_complexity_router_config(
model_params.litellm_params, deployment.litellm_params
),
model_id=_model_id,
) as table:
model_response: Final = await table.update(
where={"model_id": _model_id},
data=_data,
)
if renamed_to is not None:
await sync_access_groups_for_renamed_model(
prisma_client=prisma_client,

View file

@ -4318,6 +4318,8 @@ async def _resolve_team_access_group_resources(
access_group_id=group.access_group_id,
access_group_name=group.access_group_name,
models=tuple(group.access_model_names or ()),
mcp_server_ids=tuple(group.access_mcp_server_ids or ()),
agent_ids=tuple(group.access_agent_ids or ()),
)
for group in resolved_groups
),

View file

@ -1790,6 +1790,16 @@ def get_vertex_ai_allowed_incoming_headers(request: Request) -> dict:
return headers
def _is_vertex_anthropic_count_tokens_route(endpoint: str) -> bool:
return endpoint.rsplit("/", 1)[-1].split(":", 1)[0] == "count-tokens"
def _upstream_headers_for_vertex_route(endpoint: str, headers: Mapping[str, str]) -> Mapping[str, str]:
if not _is_vertex_anthropic_count_tokens_route(endpoint):
return headers
return MappingProxyType({name: value for name, value in headers.items() if name.lower() != "anthropic-beta"})
def get_vertex_pass_through_handler(
call_type: Literal["discovery", "aiplatform"], # noqa: UP037 # ruff reports quoted Literal values here
) -> BaseVertexAIPassThroughHandler:
@ -2188,7 +2198,7 @@ async def _base_vertex_proxy_route(
endpoint_func: Final = create_pass_through_route(
endpoint=endpoint,
target=target,
custom_headers=headers,
custom_headers=_upstream_headers_for_vertex_route(endpoint, headers),
is_streaming_request=is_streaming_request,
) # dynamically construct pass-through endpoint based on incoming path

View file

@ -40,6 +40,7 @@ from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import (
MAXIMUM_TRACEBACK_LINES_TO_LOG,
SESSION_ID_OMITTED_METADATA_KEY,
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -581,8 +582,9 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
)
# Set internal keys after merging client-supplied metadata so a request
# body that mirrors them cannot clobber the authenticated key or the
# real parent span.
# body that mirrors them cannot clobber the authenticated key, the real
# parent span, or the proxy's own session-id decision.
_metadata.pop(SESSION_ID_OMITTED_METADATA_KEY, None)
_metadata["user_api_key"] = user_api_key_dict.api_key
_metadata["litellm_parent_otel_span"] = user_api_key_dict.parent_otel_span
_metadata["user_api_key_budget_reservation"] = user_api_key_dict.budget_reservation

View file

@ -1228,6 +1228,7 @@ def run_server(
add_missing_query_params,
idle_lifetime_params,
reader_shareable_params,
translate_libpq_ssl_params,
unsupported_db_scheme,
unsupported_db_scheme_message,
)
@ -1275,11 +1276,15 @@ def run_server(
writer_url,
connection_url_params,
)
os.environ["DATABASE_URL"] = add_missing_query_params(modified_url, lifetime_params)
os.environ["DATABASE_URL"] = translate_libpq_ssl_params(
add_missing_query_params(modified_url, lifetime_params)
)
if os.getenv("DIRECT_URL", None) is not None:
database_url = os.getenv("DIRECT_URL")
modified_url = append_query_params(database_url, connection_url_params)
os.environ["DIRECT_URL"] = add_missing_query_params(modified_url, lifetime_params)
os.environ["DIRECT_URL"] = translate_libpq_ssl_params(
add_missing_query_params(modified_url, lifetime_params)
)
# The reader pool is a real pool against the same configured cap, so it
# gets the allowlisted pool params. Schema-affecting ones, including any
# the operator smuggled in through database_extra_connection_params, stay
@ -1292,14 +1297,16 @@ def run_server(
db_statement_timeout,
db_lock_timeout,
)
os.environ["DATABASE_URL_READ_REPLICA"] = add_missing_query_params(
os.environ["DATABASE_URL_READ_REPLICA"] = translate_libpq_ssl_params(
add_missing_query_params(
_with_query_value(read_replica_url, "options", reader_options)
if reader_options
else read_replica_url,
reader_shareable_params(connection_url_params),
),
lifetime_params,
add_missing_query_params(
_with_query_value(read_replica_url, "options", reader_options)
if reader_options
else read_replica_url,
reader_shareable_params(connection_url_params),
),
lifetime_params,
)
)
subprocess.run(["prisma"], capture_output=True)
is_prisma_runnable = True

View file

@ -120,6 +120,8 @@ from litellm.router_utils.add_retry_fallback_headers import (
from litellm.router_utils.auto_router_model_naming import (
STRATEGY_ROUTER_PARAM_FIELDS,
carries_complexity_router_settings,
count_heuristic_v2_routers,
heuristic_v2_limit_violation,
validate_complexity_router_config_placement,
)
from litellm.types.utils import (
@ -301,7 +303,7 @@ from litellm.proxy.auth.auth_utils import (
)
from litellm.proxy.auth.fallback_model_access import router_fallback_access_check
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.litellm_license import LicenseCheck
from litellm.proxy.auth.litellm_license import HEURISTIC_V2_LICENSE_REMEDY, LicenseCheck
from litellm.proxy.auth.model_checks import (
expand_wildcard_deployments_for_model_info,
get_all_fallbacks,
@ -4317,6 +4319,19 @@ def validate_deployment_complexity_router_placement(model: Mapping[str, object])
raise ValueError(f"model {model.get('model_name', '')!r}: {violation}")
def validate_heuristic_v2_router_limit(model_list: Sequence[Mapping[str, object]], *, limit: int | None) -> None:
"""
Refuse to start when config.yaml defines more heuristic_v2 auto-routers than the license allows.
Checked here rather than left to router registration for the same reason as the two
validators above: the proxy builds its router with `ignore_invalid_deployments=True`, so
the router's own refusal would turn the extra router into a silently missing model.
"""
violation: Final = heuristic_v2_limit_violation(held=count_heuristic_v2_routers(model_list), limit=limit)
if violation is not None:
raise ValueError(f"config.yaml model_list: {violation} {HEURISTIC_V2_LICENSE_REMEDY}")
def pin_complexity_router_model_id(model: dict) -> None: # mutable-ok: out-param, model_info is stamped in place
"""
Stamps `model_info.id` from the raw litellm_params before plugin resolution swaps
@ -5722,6 +5737,7 @@ class ProxyConfig:
model_list: Final = config.get("model_list", None)
if model_list:
router_params["model_list"] = model_list
validate_heuristic_v2_router_limit(model_list, limit=_license_check.heuristic_v2_router_limit())
print( # noqa: T201
"\033[32mLiteLLM: Proxy initialized with Config, Set models:\033[0m"
)
@ -5811,6 +5827,7 @@ class ProxyConfig:
),
ignore_invalid_deployments=True, # don't raise an error if a deployment is invalid
fallback_access_check=router_fallback_access_check,
heuristic_v2_router_limit=_license_check.heuristic_v2_router_limit,
)
if redis_usage_cache is not None and router.cache.redis_cache is None:
@ -6275,6 +6292,7 @@ class ProxyConfig:
search_tools=search_tools,
ignore_invalid_deployments=True,
fallback_access_check=router_fallback_access_check,
heuristic_v2_router_limit=_license_check.heuristic_v2_router_limit,
)
verbose_proxy_logger.debug("updated llm_router: %s", llm_router)
else:

View file

@ -4,6 +4,7 @@ import json
import os
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import (
TYPE_CHECKING,
Annotated,
@ -57,9 +58,21 @@ router: Final = APIRouter()
SPEND_LOGS_PAGINATION_COUNT_CAP: Final = 10000
_SESSION_GROUP_KEY_SQL: Final = "COALESCE(NULLIF(session_id, ''), request_id), api_key"
_SESSION_KEY_EXPR: Final = "COALESCE(NULLIF(session_id, ''), request_id)"
_SESSION_GROUP_KEY_SQL: Final = f"{_SESSION_KEY_EXPR}, api_key"
_MCP_CALL_TYPES_SQL: Final = "('call_mcp_tool', 'list_mcp_tools')"
_AGENT_CALL_TYPE_SQL: Final = "'asend_message'"
_SPEND_LOG_LIST_COLUMNS: Final = """
request_id, call_type, api_key, spend, total_tokens,
prompt_tokens, completion_tokens, "startTime", "endTime",
"completionStartTime", model, model_id, model_group,
custom_llm_provider, api_base, "user", metadata,
cache_hit, cache_key, request_tags, team_id,
organization_id, end_user, requester_ip_address,
session_id, status, mcp_namespaced_tool_name, agent_id,
COALESCE(request_duration_ms,
(EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000)::INTEGER) AS request_duration_ms
"""
_INTERNAL_HEALTH_CHECK_API_KEYS: Final = (
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
@ -160,6 +173,9 @@ class _SessionSpendRow(TypedDict):
session_cache_hit_count: ReadOnly[int]
session_llm_count: ReadOnly[int]
session_agent_count: ReadOnly[int]
session_total_prompt_tokens: ReadOnly[int]
session_total_completion_tokens: ReadOnly[int]
session_total_tokens: ReadOnly[int]
session_models: ReadOnly[Sequence[str]]
@ -175,6 +191,9 @@ class _SessionSpendStats(NamedTuple):
session_cache_hit_count: int
session_llm_count: int
session_agent_count: int
session_total_prompt_tokens: int
session_total_completion_tokens: int
session_total_tokens: int
session_models: Sequence[str]
session_models_truncated: bool
@ -2303,6 +2322,13 @@ async def ui_view_spend_logs(
default=False,
description="Paginate over sessions instead of raw logs: one representative row per session, total counts sessions",
),
session_cursor: str | None = fastapi.Query(
default=None,
description=(
"Keyset cursor '<last_activity>|<api_key>|<session_key>' from a previous group_by_session page. "
"UI route only, honored when sorting by startTime"
),
),
):
"""
View spend logs with pagination support.
@ -2636,6 +2662,18 @@ async def ui_view_spend_logs(
sql_params.append(f"%{error_message}%")
p += 1
if group_by_session is True and not is_v2 and not is_request_id_lookup and sort_by == "startTime":
return await _ui_session_grouped_spend_logs(
prisma_client=prisma_client,
sql_conditions=sql_conditions,
sql_params=sql_params,
next_param_index=p,
page=page,
page_size=page_size,
sort_desc=order_direction != "asc",
session_cursor=session_cursor,
)
# Build the ORDER BY expression. ttft_ms is computed from
# completionStartTime - startTime; non-streaming rows (where
# completionStartTime is null or equals endTime) yield NULL, so we
@ -2677,19 +2715,11 @@ async def ui_view_spend_logs(
total_is_capped: Final = raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP
total_records: Final = SPEND_LOGS_PAGINATION_COUNT_CAP if total_is_capped else raw_total
select_columns: Final = """request_id, call_type, api_key, spend, total_tokens,
prompt_tokens, completion_tokens, "startTime", "endTime",
"completionStartTime", model, model_id, model_group,
custom_llm_provider, api_base, "user", metadata,
cache_hit, cache_key, request_tags, team_id,
organization_id, end_user, requester_ip_address,
session_id, status, mcp_namespaced_tool_name, agent_id,
COALESCE(request_duration_ms, (EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000)::INTEGER) AS request_duration_ms"""
sql_query: Final = (
f"""
SELECT * FROM (
SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
{select_columns}
{_SPEND_LOG_LIST_COLUMNS}
FROM "LiteLLM_SpendLogs"
WHERE {joined_conditions}
ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
@ -2700,7 +2730,7 @@ async def ui_view_spend_logs(
if session_grouping
else f"""
SELECT
{select_columns}
{_SPEND_LOG_LIST_COLUMNS}
FROM "LiteLLM_SpendLogs"
WHERE {joined_conditions}
ORDER BY {_order_expr} {_sql_dir}{_nulls_clause}
@ -2733,6 +2763,162 @@ async def ui_view_spend_logs(
raise handle_exception_on_proxy(e)
class _SessionPageRow(TypedDict):
session_key: ReadOnly[str]
api_key: ReadOnly[str]
last_activity: ReadOnly[str]
def _parse_session_cursor(session_cursor: str | None) -> tuple[str, str, str] | None:
if session_cursor is None or session_cursor.count("|") < 2:
return None
last_activity, _, rest = session_cursor.partition("|")
api_key, _, session_key = rest.partition("|")
if not last_activity or not session_key:
return None
return (last_activity, session_key, api_key)
async def _fetch_session_representatives(
prisma_client: "PrismaClient",
where_clause: str,
sql_params: Sequence[object],
next_param_index: int,
session_keys: Sequence[tuple[str, str]],
) -> list[dict[str, object]]: # mutable-ok: _build_ui_spend_logs_response writes session counts onto each row
"""Fetch the newest non-MCP row of each ``(session_key, api_key)`` session, in ``session_keys`` order."""
rep_query: Final = f"""
SELECT * FROM (
SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
{_SPEND_LOG_LIST_COLUMNS}
FROM "LiteLLM_SpendLogs"
WHERE {where_clause}
AND ({_SESSION_GROUP_KEY_SQL}) IN (
SELECT * FROM unnest(${next_param_index}::text[], ${next_param_index + 1}::text[])
)
ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
) AS session_representatives
"""
rep_rows: Final[Sequence[dict[str, object]]] = await _query_raw( # mutable-ok: rows are enriched in place
prisma_client,
rep_query,
*sql_params,
[session_key for session_key, _ in session_keys], # mutable-ok: prisma serializes array params from a list
[api_key for _, api_key in session_keys], # mutable-ok: prisma serializes array params from a list
)
rep_by_key: Final[Mapping[tuple[str, str], dict[str, object]]] = MappingProxyType( # mutable-ok: same rows
{(str(row["session_id"] or row["request_id"]), str(row["api_key"])): row for row in rep_rows}
)
return [rep_by_key[key] for key in session_keys if key in rep_by_key] # mutable-ok: rows are enriched in place
async def _ui_session_grouped_spend_logs(
prisma_client: "PrismaClient",
sql_conditions: Sequence[str],
sql_params: Sequence[object],
next_param_index: int,
page: int,
page_size: int,
sort_desc: bool,
session_cursor: str | None,
) -> Mapping[str, object]:
"""
One row per session, keyset-paginated by session last activity.
Sessions are derived on the fly from ``LiteLLM_SpendLogs`` (no extra
table): rows sharing a ``session_id`` and ``api_key`` form a session, rows
without a session id are singletons keyed by ``request_id``. A page is the
next ``page_size`` sessions ordered by ``(MAX(startTime), session_key,
api_key)``, resumed from the ``session_cursor`` keyset
``'<last_activity>|<api_key>|<session_key>'`` instead of an OFFSET, so
page depth does not degrade the query plan. Each session is represented
by its newest non-MCP row, enriched by ``_build_ui_spend_logs_response``
exactly like the flat listing, and the response carries
``next_session_cursor`` / ``has_more`` while ``total`` counts sessions
(capped like the flat total).
"""
where_clause: Final = " AND ".join(sql_conditions) if sql_conditions else "TRUE"
cmp_op: Final = "<" if sort_desc else ">"
direction: Final = "DESC" if sort_desc else "ASC"
cursor: Final = _parse_session_cursor(session_cursor)
having_clause: Final = (
f'HAVING (MAX("startTime"), {_SESSION_GROUP_KEY_SQL}) {cmp_op} '
f"(${next_param_index}::timestamp, ${next_param_index + 1}, ${next_param_index + 2})"
if cursor
else ""
)
cursor_params: Final[tuple[object, ...]] = cursor if cursor else ()
limit_index: Final = next_param_index + len(cursor_params)
page_query: Final = f"""
SELECT {_SESSION_KEY_EXPR} AS session_key,
api_key,
MAX("startTime")::text AS last_activity
FROM "LiteLLM_SpendLogs"
WHERE {where_clause}
GROUP BY {_SESSION_GROUP_KEY_SQL}
{having_clause}
ORDER BY MAX("startTime") {direction}, {_SESSION_KEY_EXPR} {direction}, api_key {direction}
LIMIT ${limit_index}
"""
page_rows: Final[Sequence[_SessionPageRow]] = await _query_raw(
prisma_client, page_query, *sql_params, *cursor_params, page_size + 1
)
has_more: Final = len(page_rows) > page_size
visible_rows: Final = page_rows[:page_size]
next_cursor: Final = (
f"{visible_rows[-1]['last_activity']}|{visible_rows[-1]['api_key']}|{visible_rows[-1]['session_key']}"
if has_more and visible_rows
else None
)
count_query: Final = f"""
SELECT COUNT(*) AS total_count
FROM (
SELECT 1
FROM "LiteLLM_SpendLogs"
WHERE {where_clause}
GROUP BY {_SESSION_GROUP_KEY_SQL}
LIMIT ${next_param_index}
) AS bounded_sessions
"""
count_rows: Final[Sequence[_SpendLogsCountRow]] = await _query_raw(
prisma_client, count_query, *sql_params, SPEND_LOGS_PAGINATION_COUNT_CAP + 1
)
raw_total: Final = int(count_rows[0]["total_count"]) if count_rows else 0
total_is_capped: Final = raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP
total_records: Final = SPEND_LOGS_PAGINATION_COUNT_CAP if total_is_capped else raw_total
session_keys: Final = tuple((row["session_key"], row["api_key"]) for row in visible_rows)
data: Final[list[dict[str, object]]] = ( # mutable-ok: _build_ui_spend_logs_response writes onto each row
await _fetch_session_representatives(
prisma_client=prisma_client,
where_clause=where_clause,
sql_params=sql_params,
next_param_index=next_param_index,
session_keys=session_keys,
)
if session_keys
else [] # mutable-ok: downstream enrichment mutates rows in place
)
_hydrate_spend_log_metadata(data)
total_pages: Final = (total_records + page_size - 1) // page_size
response: Final[Mapping[str, object]] = await _build_ui_spend_logs_response(
prisma_client,
data,
total_records,
page,
page_size,
total_pages,
enrich_session_counts=True,
total_is_capped=total_is_capped,
)
return {**response, "next_session_cursor": next_cursor, "has_more": has_more} # mutable-ok: FastAPI response body
class RequestResponsePayload(NamedTuple):
messages: str | list | dict | None
response: str | list | dict | None
@ -4102,13 +4288,13 @@ async def _build_ui_spend_logs_response(
total_pages: int,
enrich_session_counts: bool = True,
total_is_capped: bool = False,
) -> dict:
) -> dict[str, object]:
"""
Build the paginated response for the UI spend-logs endpoint.
When ``enrich_session_counts`` is ``True`` (the default for the v1/UI
endpoint), each row is enriched with ``session_total_count`` plus spend
and call-type aggregates so the frontend knows which sessions are
endpoint), each row is enriched with ``session_total_count`` plus spend,
token and call-type aggregates so the frontend knows which sessions are
expandable (multi-call sessions). One ``GROUP BY (session_id, api_key)``
query serves every referenced session, keyed per api key so two callers
reusing a session id never see each other's totals. Rows without a
@ -4176,7 +4362,10 @@ async def _build_ui_spend_logs_response(
COUNT(*) FILTER (
WHERE call_type NOT IN {_MCP_CALL_TYPES_SQL} AND call_type != {_AGENT_CALL_TYPE_SQL}
)::int AS session_llm_count,
COUNT(*) FILTER (WHERE call_type = {_AGENT_CALL_TYPE_SQL})::int AS session_agent_count
COUNT(*) FILTER (WHERE call_type = {_AGENT_CALL_TYPE_SQL})::int AS session_agent_count,
COALESCE(SUM(prompt_tokens), 0)::bigint AS session_total_prompt_tokens,
COALESCE(SUM(completion_tokens), 0)::bigint AS session_total_completion_tokens,
COALESCE(SUM(total_tokens), 0)::bigint AS session_total_tokens
FROM "LiteLLM_SpendLogs"
WHERE session_id = ANY($1::text[])
AND api_key = ANY($2::text[])
@ -4209,6 +4398,9 @@ async def _build_ui_spend_logs_response(
session_cache_hit_count=int(row.get("session_cache_hit_count") or 0),
session_llm_count=int(row.get("session_llm_count") or 0),
session_agent_count=int(row.get("session_agent_count") or 0),
session_total_prompt_tokens=int(row.get("session_total_prompt_tokens") or 0),
session_total_completion_tokens=int(row.get("session_total_completion_tokens") or 0),
session_total_tokens=int(row.get("session_total_tokens") or 0),
session_models=models[:_SESSION_MODELS_LIMIT],
session_models_truncated=len(models) > _SESSION_MODELS_LIMIT,
)
@ -4238,6 +4430,9 @@ async def _build_ui_spend_logs_response(
row_dict["session_cache_hit_count"] = session_stats.session_cache_hit_count
row_dict["session_llm_count"] = session_stats.session_llm_count
row_dict["session_agent_count"] = session_stats.session_agent_count
row_dict["session_total_prompt_tokens"] = session_stats.session_total_prompt_tokens
row_dict["session_total_completion_tokens"] = session_stats.session_total_completion_tokens
row_dict["session_total_tokens"] = session_stats.session_total_tokens
row_dict["session_models"] = session_stats.session_models
row_dict["session_models_truncated"] = session_stats.session_models_truncated
enriched.append(row_dict)

View file

@ -15,6 +15,7 @@ from litellm.constants import (
LITELLM_TRUNCATED_PAYLOAD_FIELD,
LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE,
REDACTED_BY_LITELM_STRING,
SESSION_ID_OMITTED_METADATA_KEY,
)
from litellm.constants import (
MAX_STRING_LENGTH_PROMPT_IN_DB as DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB,
@ -578,7 +579,9 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
),
session_id=_get_session_id_for_spend_log(
kwargs=kwargs,
metadata=metadata,
standard_logging_payload=standard_logging_payload,
omit_when_missing=_omits_session_id_when_missing(metadata),
),
request_duration_ms=_get_request_duration_ms(start_time, end_time),
status=_get_status_for_spend_log(
@ -602,26 +605,39 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
raise e
def _omits_session_id_when_missing(metadata: Mapping[str, object] | None) -> bool:
"""The pre-call stamp pins `omit` on for the requests that carry it, so a config reload between pre-call and spend
logging cannot fabricate a session. `apply_missing_session_id_policy` drops any client-supplied copy of the key
from both metadata buckets before stamping, which the merge of `litellm_metadata` into `metadata` makes
necessary, so a caller cannot forge it. Requests that never reach the pre-call helper, router-model
passthrough among them, carry no stamp, so they fall back to the configured policy and `omit` still covers their
spend logs."""
if metadata is not None and metadata.get(SESSION_ID_OMITTED_METADATA_KEY):
return True
from litellm.proxy.proxy_server import general_settings
return general_settings.get("missing_session_id") == "omit"
def _get_session_id_for_spend_log(
kwargs: dict,
kwargs: Mapping[str, object],
metadata: Mapping[str, object] | None,
standard_logging_payload: StandardLoggingPayload | None,
) -> str:
"""
Get the session id for the spend log.
omit_when_missing: bool,
) -> str | None:
"""Under `omit` only `metadata.session_id`, the key Langfuse reads, counts as a session; `litellm_session_id` may
be a copied trace id."""
if omit_when_missing:
session_id: Final = metadata.get("session_id") if metadata else None
return str(session_id) if session_id else None
This ensures each spend log is associated with a unique session id.
"""
from litellm._uuid import uuid
if standard_logging_payload is not None and standard_logging_payload.get("trace_id") is not None:
return str(standard_logging_payload.get("trace_id"))
# Users can dynamically set the trace_id for each request by passing `litellm_trace_id` in kwargs
if kwargs.get("litellm_trace_id") is not None:
return str(kwargs.get("litellm_trace_id"))
# Ensure we always have a session id, if none is provided
return str(uuid.uuid4())

View file

@ -7531,6 +7531,9 @@ def create_model_info_response(
max_input_tokens = configured_input
if configured_output is not None:
max_output_tokens = configured_output
configured_mode: Final = llm_router.get_configured_mode(model_id)
if isinstance(configured_mode, str):
base["mode"] = configured_mode
if max_input_tokens is not None:
base["max_input_tokens"] = max_input_tokens

View file

@ -30,7 +30,10 @@ from litellm.proxy._types import (
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
from litellm.proxy.vector_store_endpoints.utils import can_user_access_vector_store
from litellm.proxy.vector_store_endpoints.utils import (
can_user_access_vector_store,
filter_listable_vector_stores,
)
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import ManagedVectorStoresRepository
from litellm.types.vector_stores import (
@ -390,11 +393,10 @@ async def list_vector_stores(
# Filter vector stores based on access control
accessible_vector_stores: Final = []
for vs in vector_store_map.values():
if await _check_vector_store_access(vs, user_api_key_dict):
redacted = LiteLLM_ManagedVectorStore(**vs)
redacted["litellm_params"] = _redact_sensitive_litellm_params(vs.get("litellm_params"))
accessible_vector_stores.append(redacted)
for vs in await filter_listable_vector_stores(vector_store_map.values(), user_api_key_dict):
redacted = LiteLLM_ManagedVectorStore(**vs)
redacted["litellm_params"] = _redact_sensitive_litellm_params(vs.get("litellm_params"))
accessible_vector_stores.append(redacted)
total_count: Final = len(accessible_vector_stores)
total_pages: Final = (total_count + page_size - 1) // page_size

View file

@ -1,11 +1,17 @@
import json
import re
from collections.abc import Iterable
from types import MappingProxyType
from typing import Any, Final, Literal
from fastapi import HTTPException, Request
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
is_ui_session_credential,
resolve_ui_session_team_ids,
)
from litellm.proxy._types import (
LiteLLM_ObjectPermissionTable,
LitellmUserRoles,
@ -160,10 +166,16 @@ async def can_user_access_vector_store(
if _is_proxy_admin(user_api_key_dict):
return True
vector_store_team_id: Final = vector_store.get("team_id")
if vector_store_team_id is None:
if vector_store.get("team_id") is None:
return True
return await _is_vector_store_granted(vector_store, user_api_key_dict)
async def _is_vector_store_granted(
vector_store: LiteLLM_ManagedVectorStore,
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
vector_store_id: Final = vector_store.get("vector_store_id") or ""
key_object_permission = user_api_key_dict.object_permission
@ -178,12 +190,70 @@ async def can_user_access_vector_store(
if _object_permission_allows_vector_store(team_object_permission, vector_store_id):
return True
if user_api_key_dict.team_id is not None and user_api_key_dict.team_id == vector_store_team_id:
return True
return user_api_key_dict.team_id is not None and user_api_key_dict.team_id == vector_store.get("team_id")
async def _team_auth_context(team_id: str, user_api_key_dict: UserAPIKeyAuth) -> UserAPIKeyAuth:
from litellm.proxy.auth.auth_checks import get_team_object
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
team: Final = await get_team_object(
team_id=team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_dict.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
return user_api_key_dict.model_copy(
update=MappingProxyType(
{
"team_id": team_id,
"team_object_permission": team.object_permission,
"team_object_permission_id": team.object_permission_id,
}
)
)
async def _vector_store_listing_auth_contexts(
user_api_key_dict: UserAPIKeyAuth,
) -> tuple[UserAPIKeyAuth, ...]:
if not is_ui_session_credential(user_api_key_dict):
return (user_api_key_dict,)
session_key_context: Final = user_api_key_dict.model_copy(
update=MappingProxyType({"team_id": None, "team_object_permission": None, "team_object_permission_id": None})
)
team_ids: Final = await resolve_ui_session_team_ids(user_api_key_dict)
team_contexts: Final = tuple([await _team_auth_context(team_id, user_api_key_dict) for team_id in team_ids])
return (session_key_context, *team_contexts)
async def _is_vector_store_granted_to_any(
vector_store: LiteLLM_ManagedVectorStore,
auth_contexts: tuple[UserAPIKeyAuth, ...],
) -> bool:
for auth_context in auth_contexts:
if await _is_vector_store_granted(vector_store, auth_context):
return True
return False
async def filter_listable_vector_stores(
vector_stores: Iterable[LiteLLM_ManagedVectorStore],
user_api_key_dict: UserAPIKeyAuth,
) -> tuple[LiteLLM_ManagedVectorStore, ...]:
"""Non-admins only see stores their key, one of their teams' object_permission, or team ownership grants."""
if _is_proxy_admin(user_api_key_dict):
return tuple(vector_stores)
auth_contexts: Final = await _vector_store_listing_auth_contexts(user_api_key_dict)
return tuple([vs for vs in vector_stores if await _is_vector_store_granted_to_any(vs, auth_contexts)])
async def get_litellm_managed_vector_store(
vector_store_id: str,
) -> LiteLLM_ManagedVectorStore | None:

View file

@ -8,6 +8,7 @@ from typing import Any, Final, Literal, cast
import litellm
from litellm.constants import (
AZURE_OPENAI_AUDIO_PROVIDERS,
REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
request_timeout,
@ -400,7 +401,7 @@ async def _arealtime(
litellm_metadata=_build_litellm_metadata(kwargs),
query_params=query_params,
)
elif _custom_llm_provider == "azure":
elif _custom_llm_provider in AZURE_OPENAI_AUDIO_PROVIDERS:
api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
# set API KEY
api_key = dynamic_api_key or litellm.api_key or litellm.openai_key or get_secret_str("AZURE_API_KEY")

View file

@ -562,6 +562,12 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
hidden_params: Final = getattr(chunk, "_hidden_params", None)
if hidden_params is not None:
chunk_dict["_hidden_params"] = dict(hidden_params) if isinstance(hidden_params, dict) else hidden_params
if (
chunk_dict.get("usage") is None
and isinstance(hidden_params, dict)
and hidden_params.get("usage") is not None
):
chunk_dict["usage"] = hidden_params["usage"]
return chunk_dict
def create_reasoning_summary_text_done_event(

View file

@ -21,7 +21,7 @@ import time
import traceback
import weakref
from collections import defaultdict
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Mapping, Sequence
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Generator, Iterator, Mapping, Sequence
from functools import lru_cache, partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypeVar, Union, cast
@ -117,6 +117,9 @@ from litellm.router_utils.add_retry_fallback_headers import (
from litellm.router_utils.auto_router_model_naming import (
AUTO_ROUTER_MODEL_PREFIX,
classify_strategy_router_model,
count_heuristic_v2_routers,
heuristic_v2_limit_violation,
uses_heuristic_v2_classifier,
)
from litellm.router_utils.batch_utils import (
_get_router_metadata_variable_name,
@ -211,6 +214,7 @@ from litellm.types.router import (
DeploymentTypedDict,
FallbackAccessCheck,
GuardrailTypedDict,
HeuristicV2RouterLimit,
LiteLLM_Params,
MockRouterTestingParams,
ModelGroupInfo,
@ -684,6 +688,7 @@ class Router:
background_health_check_model_groups: Sequence[str] | None = None,
enable_weighted_failover: bool = False,
fallback_access_check: FallbackAccessCheck | None = None,
heuristic_v2_router_limit: HeuristicV2RouterLimit | None = None,
) -> None:
"""
Initialize the Router class with the given parameters for caching, reliability, and routing strategy.
@ -760,6 +765,7 @@ class Router:
self.set_verbose = set_verbose
self.ignore_invalid_deployments = ignore_invalid_deployments
self.heuristic_v2_router_limit = heuristic_v2_router_limit
self.fallback_access_check: Final = fallback_access_check
self.debug_level = debug_level
self.enable_pre_call_checks = enable_pre_call_checks
@ -8794,6 +8800,30 @@ class Router:
"""
return classify_strategy_router_model(litellm_params.model) == "complexity"
def config_deployments(self) -> Iterator[Mapping[str, object]]:
"""The model_list rows that came from config.yaml rather than the DB (``model_info.db_model`` unset)."""
for deployment in self.model_list:
if not isinstance(deployment, Mapping):
continue
model_info = deployment.get("model_info")
if not (isinstance(model_info, Mapping) and model_info.get("db_model")):
yield deployment
def heuristic_v2_router_limit_violation(self) -> str | None:
"""
Why one more heuristic_v2 router cannot join this router, or None when it can.
Judged against every deployment currently on the model_list; an upsert pops the row being
edited first, so an edit of an existing heuristic_v2 router keeps its own slot. The limit is
resolved on every call through ``heuristic_v2_router_limit``; unset means unlimited, which
is the SDK default, and the proxy injects a resolver backed by its license.
"""
limit: Final = self.heuristic_v2_router_limit() if self.heuristic_v2_router_limit is not None else None
others: Final = count_heuristic_v2_routers(
deployment for deployment in self.model_list if isinstance(deployment, Mapping)
)
return heuristic_v2_limit_violation(held=others + 1, limit=limit)
def init_complexity_router_deployment(self, deployment: Deployment):
"""
Initialize the complexity-router deployment.
@ -8811,6 +8841,10 @@ class Router:
)
complexity_router_config: Final[dict | None] = deployment.litellm_params.complexity_router_config
if uses_heuristic_v2_classifier(complexity_router_config):
limit_violation: Final = self.heuristic_v2_router_limit_violation()
if limit_violation is not None:
raise ValueError(limit_violation)
default_model: str | None = deployment.litellm_params.complexity_router_default_model
@ -9634,8 +9668,16 @@ class Router:
raise e
def _restore_deployment_after_failed_upsert(self, previous_deployment: Deployment | None, model_id: str) -> None:
"""Put a deployment back the way it was before a failed upsert popped it.
A rollback re-admits state that was already serving, so it does not go through the
heuristic_v2 ceiling a newcomer gets: with the ceiling tightened since the deployment first
registered, judging the rollback would drop a serving router over an unrelated failed edit.
"""
if previous_deployment is None or self.has_model_id(model_id):
return
limit_resolver: Final = self.heuristic_v2_router_limit
self.heuristic_v2_router_limit = None
try:
self.add_deployment(deployment=previous_deployment)
verbose_router_logger.info(
@ -9650,6 +9692,8 @@ class Router:
model_id,
restore_error,
)
finally:
self.heuristic_v2_router_limit = limit_resolver
@staticmethod
def _backend_cost_map_keys(model: str, custom_llm_provider: str | None) -> tuple[str, ...]:
@ -9979,6 +10023,17 @@ class Router:
coerce_token_limit(model_info.get("max_output_tokens")),
)
def get_configured_mode(self, model_name: str) -> "str | None":
"""Return the mode explicitly configured for a concrete deployment."""
deployment: Final = self.get_deployment_by_model_group_name(model_group_name=model_name)
if deployment is None:
return None
mode: Final = deployment.model_info.get("mode")
if isinstance(mode, str) and mode.strip():
return mode
return None
def get_configured_display_name(self, model_name: str) -> "str | None":
"""
Return the display_name explicitly configured in a concrete deployment's

View file

@ -10,7 +10,7 @@ the router silently dropping the deployment at load time under
``ignore_invalid_deployments``.
"""
from collections.abc import Mapping, Sequence
from collections.abc import Iterable, Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Literal, TypeAlias
@ -163,6 +163,38 @@ def strategy_router_dependencies(
)
def uses_heuristic_v2_classifier(complexity_router_config: object) -> bool:
"""Whether this complexity config classifies with the bundled heuristic_v2 model."""
return _mapping(complexity_router_config).get("classifier_type") == "heuristic_v2"
def is_heuristic_v2_router(litellm_params: Mapping[str, object]) -> bool:
"""Whether this deployment is a complexity router that classifies with heuristic_v2."""
return classify_strategy_router_model(str(litellm_params.get("model") or "")) == "complexity" and (
uses_heuristic_v2_classifier(litellm_params.get("complexity_router_config"))
)
def count_heuristic_v2_routers(deployments: Iterable[Mapping[str, object]]) -> int:
"""How many of ``deployments`` (router model_list entries or config.yaml rows) are heuristic_v2 routers."""
return sum(1 for deployment in deployments if is_heuristic_v2_router(_mapping(deployment.get("litellm_params"))))
def heuristic_v2_limit_violation(*, held: int, limit: int | None) -> str | None:
"""Why holding ``held`` heuristic_v2 routers exceeds ``limit``, or None when it fits.
``limit`` None means unlimited. The message is shared by every enforcement point (config
load, model writes, router registration) and stays SDK-neutral: it names the cap and what
the caller can change; the proxy appends how its license lifts the cap.
"""
if limit is None or held <= limit:
return None
return (
f"At most {limit} auto-router(s) with classifier_type 'heuristic_v2' can be registered but this would make "
f"{held}. Use classifier_type 'heuristic' for this router or remove an existing heuristic_v2 router."
)
def validate_complexity_router_config_write(complexity_router_config: Mapping[str, object] | None) -> str | None:
"""Reject a complexity config the router would refuse to build a deployment from.

View file

@ -82,6 +82,7 @@ class DeploymentAffinityCheck(CustomLogger):
"""
CACHE_KEY_PREFIX = "deployment_affinity:v1"
USER_ID_AFFINITY_PREFIX: Final = "user_id:"
def __init__(
self,
@ -253,15 +254,6 @@ class DeploymentAffinityCheck(CustomLogger):
hashed_user_key: Final = cls._hash_user_key(user_key) if user_key is not None else "unscoped"
return f"{cls.CACHE_KEY_PREFIX}:session:{model_group}:{hashed_user_key}:{session_id}"
@staticmethod
def _get_user_key_from_metadata_dict(metadata: dict) -> str | None:
# NOTE: affinity is keyed on the *API key hash* provided by the proxy (not the
# OpenAI `user` parameter, which is an end-user identifier).
user_key: Final = metadata.get("user_api_key_hash")
if user_key is None:
return None
return str(user_key)
@staticmethod
def _get_session_id_from_metadata_dict(metadata: dict) -> str | None:
session_id: Final = metadata.get("session_id")
@ -285,22 +277,30 @@ class DeploymentAffinityCheck(CustomLogger):
return metadata_dicts
@staticmethod
def _get_user_key_from_request_kwargs(request_kwargs: dict) -> str | None:
def _first_metadata_value(metadata_dicts: Sequence[dict], key: str) -> str | None:
value: Final = next((metadata[key] for metadata in metadata_dicts if metadata.get(key) is not None), None)
return None if value is None else str(value)
@classmethod
def _get_user_key_from_request_kwargs(cls, request_kwargs: dict) -> str | None:
"""
Extract a stable affinity key from request kwargs.
Source (proxy): `metadata.user_api_key_hash`
Source (proxy): `metadata.user_api_key_hash` for virtual-key callers. JWT-authenticated
callers carry no key hash, so their `metadata.user_api_key_user_id` stands in for it,
namespaced under `USER_ID_AFFINITY_PREFIX` so a user id can never alias a key hash.
Note: the OpenAI `user` parameter is an end-user identifier and is intentionally
not used for deployment affinity.
"""
# Check metadata dicts (Proxy usage)
for metadata in DeploymentAffinityCheck._iter_metadata_dicts(request_kwargs):
user_key = DeploymentAffinityCheck._get_user_key_from_metadata_dict(metadata=metadata)
if user_key is not None:
return user_key
return None
metadata_dicts: Final = cls._iter_metadata_dicts(request_kwargs)
user_api_key_hash: Final = cls._first_metadata_value(metadata_dicts, "user_api_key_hash")
if user_api_key_hash is not None:
return user_api_key_hash
user_id: Final = cls._first_metadata_value(metadata_dicts, "user_api_key_user_id")
if user_id is None:
return None
return f"{cls.USER_ID_AFFINITY_PREFIX}{user_id}"
@staticmethod
def _get_session_id_from_request_kwargs(request_kwargs: dict) -> str | None:
@ -533,9 +533,9 @@ class DeploymentAffinityCheck(CustomLogger):
return typed_healthy_deployments
verbose_router_logger.debug(
"DeploymentAffinityCheck: api-key affinity hit -> deployment=%s user_key=%s",
"DeploymentAffinityCheck: caller affinity hit -> deployment=%s user_key=%s",
model_id,
self._shorten_for_logs(user_key),
self._shorten_for_logs(self._hash_user_key(user_key)),
)
return [deployment]
@ -626,7 +626,7 @@ class DeploymentAffinityCheck(CustomLogger):
deployment_model_name,
model_id,
self.ttl_seconds,
self._shorten_for_logs(user_key),
self._shorten_for_logs(self._hash_user_key(user_key)),
)
else:
verbose_router_logger.debug(

View file

@ -1,3 +1,4 @@
from collections.abc import Mapping
from typing import Any, Final
from pydantic import BaseModel, Field
@ -29,6 +30,13 @@ def is_interception_internal_key(
return any(key.startswith(prefix) for prefix in prefixes)
CONVERTED_STREAM_KEYS: Final = frozenset(f"{prefix}_converted_stream" for prefix in INTERCEPTION_INTERNAL_PREFIXES)
def converted_stream_requested(params: Mapping[str, object]) -> bool:
return any(bool(params.get(key)) for key in CONVERTED_STREAM_KEYS)
class AgenticLoopSafetyError(ValueError):
"""
Raised when an agentic-loop safety rail refuses a rerun.

View file

@ -66,6 +66,7 @@ from pydantic import (
ConfigDict,
Discriminator,
Field,
NonNegativeInt,
PrivateAttr,
SerializerFunctionWrapHandler,
field_serializer,
@ -1321,6 +1322,18 @@ class ResponseAPIUsage(BaseLiteLLMOpenAIResponseObject):
model_config = {"extra": "allow"}
class WebSearchToolUsage(BaseModel):
model_config = ConfigDict(frozen=True)
num_requests: NonNegativeInt
class ResponsesToolUsage(BaseModel):
model_config = ConfigDict(frozen=True)
web_search: WebSearchToolUsage | None = None
ResponsesAPIStatus = Literal["completed", "failed", "in_progress", "cancelled", "queued", "incomplete"]
"""
The status of the response generation.

View file

@ -11,8 +11,8 @@ class ModelInfoMetadata(TypedDict):
class ModelInfoResponse(TypedDict):
"""OpenAI-compatible model object. `mode`, `max_input_tokens`, and
`max_output_tokens` are attached when the cost map knows them; `metadata`
is present only when the endpoint is called with include_metadata=true.
`max_output_tokens` are attached when the cost map or deployment config
knows them; `metadata` is present only with include_metadata=true.
"""
id: str

View file

@ -952,6 +952,18 @@ class FallbackAccessCheck(Protocol):
async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: "Router") -> bool: ...
class HeuristicV2RouterLimit(Protocol):
"""
Resolves how many heuristic_v2 complexity routers the Router may hold right now; None means unlimited.
The Router calls it on every registration and limit query instead of caching the answer, so the
proxy can keep the limit on its license object (re-verified on config load) rather than hand
over a snapshot.
"""
def __call__(self) -> int | None: ...
class LiteLLM_RouterFileObject(TypedDict, total=False):
"""
Tracking the litellm params hash, used for mapping the file id to the right model

View file

@ -2543,6 +2543,7 @@ class ImageResponse(OpenAIImageResponse, BaseLiteLLMOpenAIResponseObject):
)
super().__init__(created=created, data=_data, usage=_usage)
self.background = kwargs.get("background", None)
self.quality = kwargs.get("quality", None)
self.output_format = kwargs.get("output_format", None)
self.size = kwargs.get("size", None)

View file

@ -3338,6 +3338,9 @@ def get_optional_params_image_gen(
continue
passed_params[k] = v
provider_supported_params: Final[tuple[str, ...]] = (
tuple(provider_config.get_supported_openai_params(model=model or "")) if provider_config is not None else ()
)
default_params: Final = {
"n": None,
"quality": None,
@ -3348,6 +3351,7 @@ def get_optional_params_image_gen(
"imageConfig": None,
"tools": None,
"web_search_options": None,
**{k: None for k in provider_supported_params},
}
non_default_params: Final = _get_non_default_params(
@ -3407,10 +3411,9 @@ def get_optional_params_image_gen(
if size is not None:
optional_params["aspectRatio"] = _map_openai_size_to_vertex_ai_aspect_ratio(size)
openai_params: list[str] = list(default_params.keys())
if provider_config is not None:
supported_params = provider_config.get_supported_openai_params(model=model or "")
openai_params = list(supported_params)
openai_params: Final[list[str]] = (
list(provider_supported_params) if provider_config is not None else list(default_params.keys())
)
optional_params = add_provider_specific_params_to_optional_params(
optional_params=optional_params,

View file

@ -29277,6 +29277,75 @@
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-6-astra": {
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
"cache_creation_input_token_cost_above_272k_tokens_flex": 1.25e-05,
"cache_creation_input_token_cost_above_272k_tokens_priority": 5e-05,
"cache_creation_input_token_cost_flex": 6.25e-06,
"cache_creation_input_token_cost_priority": 2.5e-05,
"cache_read_input_token_cost": 1e-06,
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
"cache_read_input_token_cost_above_272k_tokens_flex": 1e-06,
"cache_read_input_token_cost_above_272k_tokens_priority": 4e-06,
"cache_read_input_token_cost_flex": 5e-07,
"cache_read_input_token_cost_priority": 2e-06,
"input_cost_per_token": 1e-05,
"input_cost_per_token_above_272k_tokens": 2e-05,
"input_cost_per_token_above_272k_tokens_flex": 1e-05,
"input_cost_per_token_above_272k_tokens_priority": 4e-05,
"input_cost_per_token_batches": 5e-06,
"input_cost_per_token_flex": 5e-06,
"input_cost_per_token_priority": 2e-05,
"litellm_provider": "openai",
"max_input_tokens": 922000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5e-05,
"output_cost_per_token_above_272k_tokens": 7.5e-05,
"output_cost_per_token_above_272k_tokens_flex": 3.75e-05,
"output_cost_per_token_above_272k_tokens_priority": 0.00015,
"output_cost_per_token_batches": 2.5e-05,
"output_cost_per_token_flex": 2.5e-05,
"output_cost_per_token_priority": 0.0001,
"regional_processing_uplift_multiplier_eu": 1.1,
"regional_processing_uplift_multiplier_us": 1.1,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_cache_breakpoint": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_xhigh_reasoning_effort": true
},
"gpt-5.6": {
"cache_creation_input_token_cost": 5e-06,
"cache_creation_input_token_cost_above_272k_tokens": 1e-05,
@ -52911,6 +52980,11 @@
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
"output_cost_per_token": 3.3e-05,
"output_cost_per_token_above_272k_tokens": 4.95e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
"search_context_size_medium": 0.012
},
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@ -52933,7 +53007,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_web_search": true
},
"bedrock_mantle/openai.gpt-5.6-terra": {
"input_cost_per_token": 2.2e-06,
@ -52944,6 +53019,11 @@
"cache_read_input_token_cost_above_272k_tokens": 4.4e-07,
"output_cost_per_token": 1.32e-05,
"output_cost_per_token_above_272k_tokens": 1.98e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
"search_context_size_medium": 0.012
},
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@ -52966,7 +53046,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_web_search": true
},
"bedrock_mantle/openai.gpt-5.6-cyber": {
"input_cost_per_token": 1.375e-05,
@ -53005,6 +53086,11 @@
"cache_read_input_token_cost_above_272k_tokens": 4.4e-08,
"output_cost_per_token": 1.32e-06,
"output_cost_per_token_above_272k_tokens": 1.98e-06,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
"search_context_size_medium": 0.012
},
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@ -53027,7 +53113,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_web_search": true
},
"us.openai.gpt-5.6-sol": {
"input_cost_per_token": 4.4e-06,
@ -53192,6 +53279,11 @@
"cache_read_input_token_cost_above_272k_tokens": 1.1e-06,
"output_cost_per_token": 3.3e-05,
"output_cost_per_token_above_272k_tokens": 4.95e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
"search_context_size_medium": 0.012
},
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@ -53213,7 +53305,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_web_search": true
},
"bedrock_mantle/openai.gpt-5.4": {
"input_cost_per_token": 2.75e-06,
@ -53222,6 +53315,11 @@
"cache_read_input_token_cost_above_272k_tokens": 5.5e-07,
"output_cost_per_token": 1.65e-05,
"output_cost_per_token_above_272k_tokens": 2.475e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.012,
"search_context_size_low": 0.012,
"search_context_size_medium": 0.012
},
"litellm_provider": "bedrock_mantle",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
@ -53243,7 +53341,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_web_search": true
},
"bedrock_mantle/google.gemma-4-31b": {
"input_cost_per_token": 1.4e-07,

View file

@ -67,8 +67,8 @@ proxy = [
"azure-identity>=1.25.2,<2.0",
"azure-storage-blob>=12.28.0,<13.0",
"mcp>=1.28.1,<2.0",
"litellm-proxy-extras==0.4.92",
"litellm-enterprise==0.1.63",
"litellm-proxy-extras==0.4.93",
"litellm-enterprise==0.1.64",
"RestrictedPython>=8.5,<9.0",
"rich>=13.9.4,<14.0",
"InquirerPy>=0.3.4,<1.0",
@ -174,6 +174,8 @@ litellm-proxy = "litellm.proxy.client.cli:cli"
[dependency-groups]
dev = [
"diff-cover==9.7.2",
"hypothesis==6.165.10",
"reportlab==5.0.1",
"basedpyright==1.39.7",
"keyring==25.7.0",
"pytest==9.0.3",

View file

@ -207,13 +207,16 @@ async def completions(request: Request) -> Response:
async def embeddings(request: Request) -> Response:
body = await _parse_body(request)
model = _requested_model(body)
if model == _SLOW_MODEL:
await asyncio.sleep(_SLOW_RESPONSE_SECONDS)
raw_input = body.get("input", "")
count = len(raw_input) if isinstance(raw_input, list) else 1
return JSONResponse(
{
"object": "list",
"data": [{"object": "embedding", "index": i, "embedding": [0.0] * 1536} for i in range(max(count, 1))],
"model": _requested_model(body),
"model": model,
"usage": {"prompt_tokens": 5, "total_tokens": 5},
}
)

View file

@ -172,6 +172,7 @@ pylint: >=3.3.9 # GPLv2 license
langchain-mcp-adapters: >=0.2.1 # MIT License
langgraph: >=1.0.10 # MIT License
langgraph-prebuilt: >=1.0.8 # MIT License - https://github.com/langchain-ai/langgraph/blob/main/LICENSE
hypothesis: >=6.165.10 # MPL 2.0 license
pytest-rerunfailures: >=15.1 # MPL 2.0 license
pytest-recording: >=0.13.4 # MIT license
expression: >=5.6.0 # MIT License - https://github.com/cognitedata/Expression/blob/main/LICENSE

View file

@ -23,10 +23,12 @@ test.describe("AI Hub (internal admin view)", () => {
await expect(modal.getByText(/Select All \(\d+\)/)).toBeVisible({ timeout: 5_000 });
// Step 1: pick the seeded models via "Select All"
await modal.getByText(/Select All/i).click();
await modal.getByRole("checkbox", { name: /Select All/ }).check();
// Move to confirm step
await modal.getByRole("button", { name: "Next" }).click();
const next = modal.getByRole("button", { name: "Next" });
await expect(next).toBeEnabled();
await next.click();
await expect(modal.getByText("Confirm Making Models Public")).toBeVisible({ timeout: 5_000 });
// Submit

View file

@ -2879,7 +2879,6 @@ def response_format_tests(response: litellm.ModelResponse):
"model",
[
"bedrock/mistral.mistral-large-2407-v1:0",
"bedrock/cohere.command-r-plus-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
"mistral.mistral-7b-instruct-v0:2",
"meta.llama3-8b-instruct-v1:0",

View file

@ -15,6 +15,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import litellm
from litellm import completion, completion_cost, embedding
from tests.fake_openai_endpoint import FAKE_OPENAI_API_BASE
litellm.set_verbose = False
@ -269,11 +270,14 @@ def test_openai_azure_embedding_timeouts():
def test_openai_embedding_timeouts():
try:
response = embedding(
model="text-embedding-ada-002",
model="openai/slow-endpoint",
input=["good morning from litellm"],
timeout=0.00001,
api_base=FAKE_OPENAI_API_BASE,
api_key="fake-key",
timeout=0.5,
)
print(response)
pytest.fail("Expected timeout error, the request returned instead")
except openai.APITimeoutError:
print("Good job got OpenAI timeout error!")
pass

View file

@ -1552,8 +1552,9 @@ def test_router_timeout():
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": "os.environ/OPENAI_API_KEY",
"model": "openai/slow-endpoint",
"api_base": FAKE_OPENAI_API_BASE,
"api_key": "fake-key",
},
}
]
@ -1562,7 +1563,7 @@ def test_router_timeout():
start_time = time.time()
try:
res = router.completion(
model="gpt-3.5-turbo", messages=messages, timeout=0.0001
model="gpt-3.5-turbo", messages=messages, timeout=0.5
)
print(res)
pytest.fail("this should have timed out")

View file

@ -1168,7 +1168,6 @@ async def test_completion_replicate_llama3_streaming(sync_mode):
"model, region",
[
# ["bedrock/ai21.jamba-instruct-v1:0", "us-east-1"],
# ["bedrock/cohere.command-r-plus-v1:0", None],
["us.anthropic.claude-sonnet-4-5-20250929-v1:0", None],
# ["mistral.mistral-7b-instruct-v0:2", None],
# ["meta.llama3-8b-instruct-v1:0", None],
@ -1271,7 +1270,7 @@ def test_bedrock_claude_3_streaming():
"model",
[
"claude-haiku-4-5-20251001",
"cohere.command-r-plus-v1:0", # bedrock
"bedrock/mistral.mistral-7b-instruct-v0:2",
"gpt-3.5-turbo",
],
)

View file

@ -12,6 +12,7 @@ import openai
import pytest
import litellm
from tests.fake_openai_endpoint import FAKE_OPENAI_API_BASE
@pytest.mark.parametrize(
@ -216,13 +217,16 @@ def test_timeout_streaming():
litellm.set_verbose = False
try:
response = litellm.completion(
model="gpt-3.5-turbo",
model="openai/slow-endpoint",
messages=[{"role": "user", "content": "hello, write a 20 pg essay"}],
timeout=0.0001,
api_base=FAKE_OPENAI_API_BASE,
api_key="fake-key",
timeout=0.5,
stream=True,
)
for chunk in response:
print(chunk)
pytest.fail("Did not raise error `openai.APITimeoutError`. The stream completed instead")
except openai.APITimeoutError as e:
print(
"Passed: Raised correct exception. Got openai.APITimeoutError\nGood Job", e

View file

@ -2815,6 +2815,7 @@ async def test_mcp_server_manager_with_access_groups_integration():
"""Integration test for MCPServerManager with access group filtering"""
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
MCPServerAccess,
)
from litellm.proxy._types import UserAPIKeyAuth
@ -2848,11 +2849,11 @@ async def test_mcp_server_manager_with_access_groups_integration():
)
# Mock the permission lookup to return staff access group
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers") as mock_get_allowed:
mock_get_allowed.return_value = [
"staff-server-id",
"ops-server-id",
] # User has access to staff and ops
with patch.object(MCPRequestHandler, "get_mcp_server_access") as mock_get_allowed: # test-quality-ok: manager resolver seam
mock_get_allowed.return_value = MCPServerAccess(
server_ids=("staff-server-id", "ops-server-id"),
scope="scoped",
)
allowed_servers = await test_manager.get_allowed_mcp_servers(user_auth)
@ -2901,6 +2902,7 @@ async def test_get_allowed_mcp_servers_returns_empty_for_non_admin_without_permi
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
MCPServerAccess,
)
test_manager = MCPServerManager()
@ -2923,9 +2925,9 @@ async def test_get_allowed_mcp_servers_returns_empty_for_non_admin_without_permi
)
with patch.object(
MCPRequestHandler, "get_allowed_mcp_servers", new_callable=AsyncMock
MCPRequestHandler, "get_mcp_server_access", new_callable=AsyncMock
) as mock_permission_lookup:
mock_permission_lookup.return_value = []
mock_permission_lookup.return_value = MCPServerAccess(server_ids=())
allowed_servers = await test_manager.get_allowed_mcp_servers(user_auth)
assert allowed_servers == []

View file

@ -1,148 +1,105 @@
# Rust ↔ Python SDK parity harness
# Rust/Python migration harness
This folder is the operator-facing harness for the Rust migration test plan. It runs pytest normally, listens to test events in-process, and redraws a live matrix grouped by testing strategy and SDK-level function.
This local harness follows [the agreed structure](AGENTS.md). The root command selects strategies and combines their reports. Each strategy has an independent entry point
The matrix always has these SDK columns:
- `ocr / aocr`
- `messages / amessages`
- `responses / aresponses`
- `count_tokens`
- `chat_completions / acompletion`
- `transcription / atranscription`
The harness has four deliberately broad test-strategy folders:
| Strategy | Folder |
| --- | --- |
| Public SDK parity over generated and recorded inputs | [`e2e_fuzz_tests/`](e2e_fuzz_tests/) |
| Focused tests of Rust-owned behavior | [`unit_tests_rust/`](unit_tests_rust/) |
| Isolated transform and Python-to-Rust helper coverage | [`validate_sub_methods/`](validate_sub_methods/) |
| Already-existing live-API SDK tests | [`existing_e2e_test_sdk/`](existing_e2e_test_sdk/) |
## Run it
From the repository root:
```bash
poetry run python -m tests.rust-python-harness
```text
strategies/
e2e_parity/runner.py
sdk/ocr/fixtures/
sdk/messages/
sdk/chat_completions/
sdk/responses/
gateway/
existing_e2e_test_sdk/runner.py
trace_parity/runner.py
sdk/
gateway/
unit_tests/
runner.py
mapping_validator.py
python_runner.py
rust_runner.py
shared/
parity/
tracing/
reporting/
```
The default runs every configured test once and updates all matching cells in real time. Narrow a run by strategy, SDK function, or both:
## Run locally
```bash
poetry run python -m tests.rust-python-harness --strategy e2e_fuzz_tests
poetry run python -m tests.rust-python-harness --function messages
poetry run python -m tests.rust-python-harness --strategy validate_sub_methods --function ocr
uv run python -m tests.rust-python-harness --list
uv run python -m tests.rust-python-harness --function ocr --plain
uv run python -m tests.rust-python-harness --strategy e2e_parity --surface sdk --function ocr --plain
uv run python -m tests.rust-python-harness.strategies.e2e_parity.runner --function ocr --plain
uv run python -m tests.rust-python-harness.strategies.trace_parity.runner --plain
uv run python -m tests.rust-python-harness.strategies.unit_tests.runner --plain
uv run python -m tests.rust-python-harness.strategies.existing_e2e_test_sdk.runner --function transcription --plain
```
For a guided run, use the interactive picker. It asks which strategy rows and SDK
function columns to include, then hands the terminal to the live dashboard. It never
captures keys while tests are running, so Ctrl-C and pytest debugging remain safe.
Use `--interactive` for strategy and function selection, `--pytest-arg=-x` to stop pytest on its first failure, and `--coverage` to write Python coverage under `target/rust-python-harness/`. The harness enables pytest namespace-package discovery only for its own invocations
```bash
poetry run python -m tests.rust-python-harness --interactive
```
This harness has no CI execution. A configured test that fails or disappears makes the command fail. An unconfigured strategy cell remains planned and contributes no passing evidence. Interruptions and collection errors stop execution; ordinary test failures remain in the combined report while later strategies run
Useful operator options:
## Strategy responsibilities
```bash
# Inspect coverage and pytest selectors without running anything.
poetry run python -m tests.rust-python-harness --list
E2E parity compares SDK objects, exceptions, callbacks, streams, and provider requests. Gateway tests compare HTTP responses. Both surfaces use the same strategy runner and keep execution details and fixtures in their own folders. OCR has recorded sync/async SDK coverage; the existing Messages and Responses bridge checks remain partial
# Stable line-oriented output for CI logs or redirected output.
poetry run python -m tests.rust-python-harness --plain
Trace parity compares operation names through an explicit Python/Rust mapping, call counts, and required completion-before-start ordering with `shared/tracing/compare.py`. Surface tests supply captured operation intervals. No production trace instrumentation or trace case is configured yet
# Measure Python reference lines exercised by this parity run and build an HTML heatmap.
poetry run python -m tests.rust-python-harness --coverage
Unit testing combines test mapping validation, separate Python processes with Rust disabled and enabled, backend verification, result comparison, and native Cargo tests. Native tests stay beside their Rust implementation. Existing Python tests stay at their original paths. No complete Python/native unit mapping is configured yet, so these cells remain planned
# Forward pytest options. Use the equals form when the value begins with a dash.
poetry run python -m tests.rust-python-harness --pytest-arg=-x
```
The existing E2E SDK strategy retains the live provider tests configured upstream. It runs OCR, Chat Completions, and Transcription checks from their existing paths and reports them separately from parity tests. These tests require provider credentials
The process returns pytest's exit code. A configured selector that collects no test is also a failure. A planned cell has no selector yet and does not fail the run.
## Configure cases
The dashboard adapts to narrow terminals, shows elapsed time and unique-test progress,
and prints the three slowest tests when the run ends. Each failure includes a focused
`poetry run pytest ... -q` command. Redirected output and CI automatically use the
line-oriented plain renderer; `--plain` lets you opt into it locally.
The final screen includes a confidence score for every SDK section. It is the direct
ratio of required strategy rows with passing evidence, such as `1/3 = 33%`; High means
all required strategies passed, Medium means some passed, and Low means none passed.
This behavioral score is intentionally shown separately from Python and Rust LOC.
Coverage reports are written outside the three strategy folders at
`target/rust-python-harness/`. Open `python-html/index.html` to inspect executed and
missing Python lines; `python.json` and `python.xml` are available for automation.
Coverage is finalized after pytest exits, because worker processes must flush their
data first.
## Port coverage and confidence
Treat these as separate signals instead of one ambiguous coverage percentage:
| Signal | Tool | What it proves |
| --- | --- | --- |
| Python reference LOC | `coverage.py` / `pytest-cov` via `--coverage` | The mapped Python behavior ran |
| Rust port LOC | `cargo-llvm-cov` | The mapped Rust implementation ran |
| Parity contracts | This harness matrix | Python and Rust had the same observable behavior |
`validate_sub_methods/` owns the future source-section inventory that maps a stable
Python qualified symbol to its Rust symbol. That inventory is the denominator for
per-function rollups; raw coverage for the entire LiteLLM repository would obscure
the port's real gaps. `unit_tests_rust/` owns direct `cargo-llvm-cov` runs, while
`e2e_fuzz_tests/` owns behavioral parity and fuzz-case counts. Keep Python, Rust, and
parity percentages visible side by side and label section confidence High only when
the mapped implementation exists, every required strategy passes, and both sides meet
their LOC thresholds. Generated Rust LCOV/HTML and the combined index also belong in
`target/rust-python-harness/`, not in a fourth strategy folder.
## Read the matrix
| Mark | Meaning |
| --- | --- |
| `✓` | All collected tests passed |
| `✗` | At least one test failed |
| `!` | Test setup or teardown failed |
| `↷` | All collected tests skipped |
| `?` | A configured selector did not collect a test |
| `—` | Strategy is planned but has no test yet |
| `n/a` | Strategy does not apply to this SDK function |
| `◐` | The configured tests cover only part of the TDD's parity contract |
The initial end-to-end entries deliberately show `◐`: the repository has Rust bridge tests for OCR, Messages, and Responses websocket plumbing, but those are not yet frozen-Python-oracle comparisons. The remaining TDD cells stay visible as planned work instead of disappearing from a green summary.
## Attach parity tests
Each of the four folders contains a concise `README.md` and a `strategy.json`. Add a pytest file or node ID to the appropriate SDK function's `selectors` list:
Each strategy has a `strategy.json`. Its `functions` object defines SDK cases for OCR, Messages, Responses, Count Tokens, Chat Completions, and Transcription. E2E and trace manifests also accept a `gateway` object keyed by API name. A case has `coverage`, `selectors`, and an optional `note`
```json
{
"coverage": "complete",
"selectors": [
"tests/rust-python-harness/validate_sub_methods/test_messages.py"
]
"coverage": "partial",
"selectors": ["tests/rust-python-harness/strategies/e2e_parity/sdk/ocr/test_sdk_parity.py"]
}
```
Selectors use the same syntax as pytest. A file selector aggregates every test in the file; a node selector can target one test or parametrized family; a selector ending in `/` aggregates every test in that folder, recursively. The runner deduplicates selectors, so one test may intentionally prove more than one cell without executing twice.
Selectors use pytest file or node syntax. A selector ending in `/` includes tests recursively from that directory
Use these coverage values:
Use `planned` with no selectors until an executable contract exists, `partial` for incomplete coverage, `complete` for the full contract, and `not_applicable` when a strategy does not apply. The dashboard shows passing evidence separately from coverage completeness and LOC coverage
- `complete`: implements the full strategy contract for that SDK function.
- `partial`: useful coverage exists, but the TDD contract is not fully proven.
- `planned`: no runnable parity test exists yet.
- `not_applicable`: the strategy cannot apply, such as streaming for OCR.
Unit cases use `unit_suite` instead of `selectors`, pointing to a repository-relative JSON file with this shape:
Keep comparison mechanics in shared harness modules and provider/function facts in the owning strategy folder. A Python/Rust mismatch is a test failure; do not normalize away observable return types, exception classes, private response fields, chunk ordering, or callback payload differences merely to make a cell green.
```json
{
"python_selectors": ["tests/test_api.py::test_decode"],
"cargo_manifest": "litellm-rust/Cargo.toml",
"cargo_package": "litellm-core",
"cargo_filter": "ocr::",
"backend": {
"environment_variable": "LITELLM_USE_RUST_OCR",
"probe": "tests.rust-python-harness.strategies.unit_tests.python_runner:ocr_backend"
},
"mappings": [{"python": "tests/test_api.py::test_decode", "rust": "ocr::test_decode"}]
}
```
## Architecture
Names match automatically when the collected Python and Rust test names agree. Explicit `mappings` handle different names, class names, and parametrized cases. Missing or ambiguous counterparts fail validation in either direction. The Cargo filter must select the same behavior as the Python selectors
- `catalog.py` validates and loads every strategy manifest.
- `models.py` owns typed strategy, case, coverage, and run-state models.
- `runner.py` maps live pytest events back to one or more matrix cells.
- `ui.py` renders the interactive Rich dashboard and a dependency-free plain fallback.
- `cli.py` handles filtering and preserves pytest exit semantics.
The backend probe returns `python` or `rust` and runs at startup and before every test call, after fixtures have run. The OCR probe verifies the dispatch flag and native extension availability. Surface tests must also assert that calls reach their intended implementation to catch per-call fallback. Python outcomes must agree, and failed runs remain failures even if both backends fail identically
The harness is driven from Python, matching the SDK surface and existing test tooling. Rust remains responsible for the implementation under comparison; the harness does not move provider semantics into the PyO3 bridge.
## OCR fixtures
Fixtures, provider configuration, input strategies, and recording commands live in [the OCR package](strategies/e2e_parity/sdk/ocr/fixtures/README.md). Record with provider credentials:
```bash
uv run python -m tests.rust-python-harness.strategies.e2e_parity.sdk.ocr.fixtures.record --examples 1000
```
`LITELLM_OCR_FIXTURE_DIR` and `--fixture-dir` override the default directory. Shared recording, replay, comparison, streaming, and cassette persistence live in `shared/parity/`
Run the harness's own checks locally:
```bash
uv run pytest -o consider_namespace_packages=true tests/rust-python-harness/shared tests/rust-python-harness/strategies/unit_tests tests/test_rust_python_harness.py -q
```
Existing OCR parity gaps remain visible: invalid-model provider errors differ, Reducto lacks a native contract, and the expanded Azure corpus exposes duplicate Content-Type headers. Moving the harness does not change provider responses or weaken assertions

View file

@ -2,92 +2,74 @@ from __future__ import annotations
import json
from pathlib import Path
from typing import Any
from typing import Final
from .models import Coverage, HarnessCase, SDK_FUNCTIONS, Strategy
from pydantic import BaseModel, ConfigDict, ValidationError
STRATEGIES_ROOT = Path(__file__).parent
from .shared.reporting.models import Coverage, HarnessCase, SDK_FUNCTIONS, Strategy
STRATEGIES_ROOT: Final = Path(__file__).parent / "strategies"
def _require_string(value: Any, field: str, source: Path) -> str:
if not isinstance(value, str) or not value.strip():
raise ValueError(f"{source}: {field} must be a non-empty string")
return value
class CaseSpec(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
coverage: Coverage
selectors: tuple[str, ...] = ()
note: str = ""
unit_suite: str | None = None
class StrategySpec(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
order: int
id: str
label: str
description: str
functions: dict[str, CaseSpec]
gateway: dict[str, CaseSpec] = {}
def _load_strategy(source: Path) -> Strategy:
with source.open(encoding="utf-8") as stream:
data = json.load(stream)
strategy_id = _require_string(data.get("id"), "id", source)
label = _require_string(data.get("label"), "label", source)
description = _require_string(data.get("description"), "description", source)
order = data.get("order")
if not isinstance(order, int):
raise ValueError(f"{source}: order must be an integer")
function_data = data.get("functions")
if not isinstance(function_data, dict):
raise ValueError(f"{source}: functions must be an object")
missing = set(SDK_FUNCTIONS) - set(function_data)
extra = set(function_data) - set(SDK_FUNCTIONS)
if missing or extra:
raise ValueError(
f"{source}: functions must exactly match {SDK_FUNCTIONS}; missing={missing}, extra={extra}"
data: Final = StrategySpec.model_validate_json(source.read_text(encoding="utf-8"))
if set(data.functions) != set(SDK_FUNCTIONS):
raise ValueError(f"{source}: functions must exactly match {SDK_FUNCTIONS}")
cases: Final = tuple(
HarnessCase(
strategy_id=data.id,
strategy_label=data.label,
sdk_function=name,
coverage=case.coverage,
selectors=case.selectors,
note=case.note,
surface=surface,
unit_suite=case.unit_suite,
)
cases: list[HarnessCase] = []
for sdk_function in SDK_FUNCTIONS:
case_data = function_data[sdk_function]
if not isinstance(case_data, dict):
raise ValueError(f"{source}: functions.{sdk_function} must be an object")
try:
coverage = Coverage(case_data.get("coverage"))
except ValueError as exc:
raise ValueError(f"{source}: invalid coverage for {sdk_function}") from exc
selectors = case_data.get("selectors", [])
if not isinstance(selectors, list) or not all(
isinstance(item, str) and item for item in selectors
):
raise ValueError(
f"{source}: selectors for {sdk_function} must be a list of strings"
)
if coverage is Coverage.NOT_APPLICABLE and selectors:
raise ValueError(
f"{source}: not_applicable case {sdk_function} cannot have selectors"
)
cases.append(
HarnessCase(
strategy_id=strategy_id,
strategy_label=label,
sdk_function=sdk_function,
coverage=coverage,
selectors=tuple(selectors),
note=str(case_data.get("note", "")),
)
)
return Strategy(
order=order,
id=strategy_id,
label=label,
description=description,
directory=source.parent,
cases=tuple(cases),
for surface, functions in (("sdk", data.functions), ("gateway", data.gateway))
for name in (SDK_FUNCTIONS if surface == "sdk" else functions)
for case in (functions[name],)
)
for case in cases:
if case.coverage in {Coverage.PLANNED, Coverage.NOT_APPLICABLE} and (case.selectors or case.unit_suite):
raise ValueError(f"{source}: {case.coverage.value} case {case.key} cannot configure tests")
if any(not selector.strip() for selector in case.selectors):
raise ValueError(f"{source}: empty selector in {case.key}")
if data.id == "unit_tests" and case.selectors:
raise ValueError(f"{source}: unit_tests must configure unit_suite instead of pytest selectors")
if data.id != "unit_tests" and case.unit_suite:
raise ValueError(f"{source}: unit_suite is only valid for unit_tests")
return Strategy(data.order, data.id, data.label, data.description, source.parent, cases)
def load_catalog(root: Path = STRATEGIES_ROOT) -> tuple[Strategy, ...]:
sources = sorted(root.glob("*/strategy.json"))
sources: Final = tuple(sorted(root.glob("*/strategy.json")))
if not sources:
raise ValueError(f"No strategy manifests found below {root}")
strategies = tuple(
sorted(
(_load_strategy(source) for source in sources),
key=lambda strategy: strategy.order,
)
)
ids = [strategy.id for strategy in strategies]
if len(ids) != len(set(ids)):
try:
strategies: Final = tuple(sorted((_load_strategy(source) for source in sources), key=lambda item: item.order))
except (ValidationError, json.JSONDecodeError) as error:
raise ValueError(str(error)) from error
if len({strategy.id for strategy in strategies}) != len(strategies):
raise ValueError(f"Duplicate strategy id in {root}")
return strategies

View file

@ -6,10 +6,14 @@ from collections.abc import Sequence
from pathlib import Path
from .catalog import load_catalog
from .models import SDK_FUNCTIONS, HarnessCase, Strategy
from .runner import run_pytest
from .ui import make_dashboard
from .shared.reporting.models import SDK_FUNCTIONS, HarnessCase, Strategy
from .shared.reporting.orchestration import StrategyRunner, run_strategies
from .shared.reporting.ui import make_dashboard
from .strategies.e2e_parity.runner import run as run_e2e
from .strategies.existing_e2e_test_sdk.runner import run as run_existing
from .strategies.trace_parity.runner import run as run_trace
from .strategies.unit_tests.mapping_validator import FunctionReport, build_function_report
from .strategies.unit_tests.runner import run as run_units
REPO_ROOT = Path(__file__).resolve().parents[2]
COVERAGE_ROOT = REPO_ROOT / "target" / "rust-python-harness"
@ -44,6 +48,7 @@ def _parser() -> argparse.ArgumentParser:
choices=SDK_FUNCTIONS,
help="run only this SDK function",
)
parser.add_argument("--surface", choices=("sdk", "gateway"), help="run only this API surface")
parser.add_argument(
"--validate-ledger",
action="store_true",
@ -135,9 +140,9 @@ def _print_catalog(strategies: Sequence[Strategy]) -> None:
print(f"{strategy.id:20} {strategy.label}")
for case in strategy.cases:
selectors = (
", ".join(case.selectors) if case.selectors else "no test configured"
", ".join(case.selectors) if case.selectors else case.unit_suite or "no test configured"
)
print(f" {case.sdk_function:12} {case.coverage.value:14} {selectors}")
print(f" {case.surface}/{case.sdk_function:12} {case.coverage.value:14} {selectors}")
def _print_function_report(report: FunctionReport) -> None:
@ -172,7 +177,21 @@ def _validate_ledger(sdk_functions: set[str]) -> int:
return 0 if all(report.is_clean for report in reports) else 1
def main(argv: Sequence[str] | None = None) -> int:
def _resolve_runner(strategy_id: str) -> StrategyRunner:
match strategy_id:
case "e2e_parity":
return run_e2e
case "trace_parity":
return run_trace
case "unit_tests":
return run_units
case "existing_e2e_test_sdk":
return run_existing
case _:
raise ValueError(f"Unknown strategy: {strategy_id}")
def main(argv: Sequence[str] | None = None, *, strategy_id: str | None = None) -> int:
args = _parser().parse_args(argv)
if args.coverage and importlib.util.find_spec("pytest_cov") is None:
_parser().error(
@ -181,7 +200,8 @@ def main(argv: Sequence[str] | None = None) -> int:
)
if args.validate_ledger:
return _validate_ledger(set(args.sdk_functions))
strategies = load_catalog()
catalog = load_catalog()
strategies = tuple(strategy for strategy in catalog if strategy_id is None or strategy.id == strategy_id)
if args.list:
_print_catalog(strategies)
return 0
@ -194,7 +214,8 @@ def main(argv: Sequence[str] | None = None) -> int:
sdk_functions = sdk_functions or picked_functions
try:
cases = _select(strategies, strategy_ids, sdk_functions)
selected = _select(strategies, strategy_ids, sdk_functions)
cases = tuple(case for case in selected if args.surface is None or case.surface == args.surface)
except ValueError as exc:
_parser().error(str(exc))
selected_strategy_ids = {case.strategy_id for case in cases}
@ -210,11 +231,12 @@ def main(argv: Sequence[str] | None = None) -> int:
if args.coverage:
pytest_args.extend(_coverage_pytest_args())
with dashboard:
exit_code, run = run_pytest(
exit_code, run = run_strategies(
cases=cases,
repo_root=REPO_ROOT,
on_update=dashboard.update,
pytest_args=pytest_args,
resolve_runner=_resolve_runner,
)
dashboard.finish(run, exit_code)
if args.coverage and (COVERAGE_ROOT / "python.json").exists():

View file

@ -1,3 +0,0 @@
# End-to-end fuzz tests
Runs the same SDK call through the Python and Rust paths using generated inputs and recorded provider responses. It compares public results, streams, callbacks, and exceptions to catch behavior differences a unit test can miss.

View file

@ -1,14 +0,0 @@
{
"order": 10,
"id": "e2e_fuzz_tests",
"label": "End-to-end fuzz tests",
"description": "Compare observable Python and Rust SDK behavior over generated and recorded inputs.",
"functions": {
"ocr": {"coverage": "partial", "selectors": ["tests/test_litellm/ocr/test_rust_bridge.py"], "note": "Bridge coverage exists; frozen-oracle fuzz parity is still being added."},
"messages": {"coverage": "partial", "selectors": ["tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py"], "note": "Bridge coverage exists; frozen-oracle fuzz parity is still being added."},
"responses": {"coverage": "partial", "selectors": ["tests/test_litellm/responses/test_rust_bridge_websocket.py"], "note": "Covers the websocket bridge; full responses parity is still being added."},
"count_tokens": {"coverage": "planned", "selectors": [], "note": "No Rust count_tokens parity test is present yet."},
"chat_completions": {"coverage": "partial", "selectors": ["tests/test_litellm/rust_bridge/test_chat_completions.py"], "note": "Bridge coverage exists; frozen-oracle fuzz parity is still being added."},
"transcription": {"coverage": "partial", "selectors": ["tests/test_litellm/test_audio_transcription_rust_bridge.py"], "note": "Bridge coverage exists; frozen-oracle fuzz parity is still being added."}
}
}

View file

@ -0,0 +1,91 @@
# Implementation parity testing through the SDK interface
> Given the same SDK call and identical provider behavior, do two implementations expose the same SDK contract?
## What the harness compares
- A fixture contains a LiteLLM SDK input and a recorded upstream provider response
- The same LiteLLM input is transformed by isolated baseline and candidate implementations
- The resulting provider requests must match in method, path, headers, and body, excluding runtime-specific HTTP metadata
- The recorded provider response is then replayed unchanged to both workers
- The harness compares the values returned through the Python SDK interface
- Non-streaming responses are compared directly, including their concrete return type and public model fields
- Streaming responses are consumed and compared chunk by chunk, including wrapper type, chunk type and order, termination, and public exception behavior
- Failed SDK calls are compared by exception class, stable message, status, code, model, provider, and parameter fields
- Traceback paths and line numbers are excluded because they are runtime-specific
- Route-specific comparators and chunk normalizers handle differences in each public SDK contract
## Process isolation
- SDK object and stream parity runs both implementations sequentially in the same process so tests can retain returned objects
- Every test saves and restores the original bridge state
- A small subprocess smoke test verifies environment-based startup configuration and detects fallback to the Python HTTP implementation
## Streaming execution
The invocation callback passed to `run_in_process` must consume the stream before returning its `StreamOutcome`.
Use `consume_sync_stream` inside that callback, or await `consume_async_stream` inside the callback passed to
`run_in_process_async`. Provider requests are collected only after the callback completes. Streaming is explicit:
an iterable return value alone does not select stream consumption
The consumers retain the wrapper type, iteration capabilities, chunk types and order, and any partial output before
an error. Errors retain their creation or iteration phase and the full public `SDKError` fields, with traceback text
removed. `capture_sync_stream` and `capture_async_stream` consume through the same helpers and then serialize the
outcome for subprocess reports. A serialization failure raises as a harness failure rather than becoming an SDK error
Response models and stream chunks share a recursive comparator. It compares concrete model, container, and scalar
types, public fields and extras, and exact values while ignoring Pydantic private attributes at every nesting level.
An API may supply an explicit chunk normalizer for its public contract
Shared tests exercise a local SSE provider through recording, VCR cassette storage, replay, and typed event comparison
in sync and async modes. They cover fragmented events, split UTF-8 characters, CRLF framing, coalesced events, and
application errors within a normally completed HTTP stream. HTTP byte boundaries and decoded SDK event boundaries
are checked separately
OCR remains the only integrated LiteLLM route. These tests validate shared streaming machinery, not another route's
SDK parity. Connection interruption, early cancellation, and lifecycle timeout enforcement remain outside this coverage
## Hypothesis and property-based testing
- Hypothesis is a Python library for property-based testing
- Example-based tests use inputs selected by the test author
- Property-based tests define strategies for valid inputs and properties that must hold for every generated example
- Hypothesis generates combinations from those strategies and normally shrinks a failing example to a smaller reproducible case
- In this harness, Hypothesis is used only during fixture generation to expand the LiteLLM input corpus
- Each API owns the strategies that vary its supported inputs
- Fixture generation is deterministic, and each generated input is recorded with the raw provider response it received
- The parity tests use committed fixtures and do not call the provider or generate new Hypothesis examples
- Provider responses are replayed unchanged, so the parity test does not fuzz or validate provider behavior
- Because Hypothesis does not run the parity assertion directly, parity failures are not automatically shrunk
## API-owned fixtures
The shared package owns recording, replay, persistence, execution, comparison, and route-neutral media constructors.
Each API package owns its input models, explicit strategies, provider targets, route-specific assets, fixture directory,
and regeneration command. See the API package documentation for its configured contracts and recording command
## VCR cassettes
Fixtures use VCR's YAML `version: 1` format with ordered request/response `interactions`. VCR handles text and binary
body serialization. Each cassette also contains `recorded_at`, `ttl_seconds: 0` (committed fixtures never expire), and
`x-litellm` metadata holding the SDK input and request provenance. Streaming responses carry
`x-litellm-chunk-lengths` so local replay preserves the original byte boundaries
The recording server captures requests before forwarding their responses. Saved requests use the stable
`http://parity-provider.invalid` origin and strip authentication headers and credential query parameters. The upstream
request keeps its credentials. Provider response bytes and non-success statuses are preserved
Standard VCR can load these files and replay their interactions. Parity tests keep using the local HTTP server because
Rust HTTP calls do not pass through VCR's Python patches. The harness still compares the two implementations' requests
against each other; the saved request is available for inspection and VCR playback, not a new parity assertion
Refresh parity cassettes through the API's recording command. Generic VCR writers do not preserve the SDK metadata
Legacy JSON fixtures remain readable. Migrated cassettes mark reconstructed requests as `python_replay`; fresh
recordings use `recorded`. The metadata extensions follow the filesystem cassette layout proposed in
[PR #39338](https://github.com/BerriAI/litellm/pull/39338), without depending on its unmerged persistence backend
## References
- [Hypothesis documentation](https://hypothesis.readthedocs.io/en/latest/)
- [Hypothesis documentation source](https://github.com/HypothesisWorks/hypothesis/tree/master/hypothesis/docs)

View file

@ -0,0 +1,3 @@
import pytest
pytest.register_assert_rewrite("tests.rust-python-harness.shared.parity.compare")

View file

@ -0,0 +1,80 @@
from __future__ import annotations
from collections.abc import Mapping, Sequence
from typing import Final, cast
from pydantic import BaseModel
from .models import CapturedRequest, Execution
def validate_harness(baseline: Execution, candidate: Execution, baseline_user_agent: str) -> None:
for request in baseline.requests:
if request.user_agent != baseline_user_agent:
raise AssertionError(
f"baseline provider request did not carry sentinel user-agent {baseline_user_agent!r}: "
f"{request.user_agent!r}"
)
for request in candidate.requests:
if request.user_agent == baseline_user_agent:
raise AssertionError("candidate route fell back to the baseline HTTP implementation")
def _request_after_transformation(request: CapturedRequest) -> CapturedRequest:
return request.model_copy(update={"user_agent": None})
def assert_request_parity(baseline: tuple[CapturedRequest, ...], candidate: tuple[CapturedRequest, ...]) -> None:
baseline_requests: Final = tuple(_request_after_transformation(request) for request in baseline)
candidate_requests: Final = tuple(_request_after_transformation(request) for request in candidate)
assert_value_parity(baseline_requests, candidate_requests)
def _public_model_values(model: BaseModel) -> dict[str, object]:
fields: Final = (*type(model).model_fields, *type(model).model_computed_fields)
extras: Final = cast(Mapping[str, object], model.model_extra or {})
return {
**{name: cast(object, getattr(model, name)) for name in fields if not name.startswith("_")},
**{name: value for name, value in extras.items() if not name.startswith("_")},
}
def assert_model_parity(baseline: BaseModel, candidate: BaseModel) -> None:
assert_value_parity(baseline, candidate)
def assert_value_parity(baseline: object, candidate: object, *, path: str = "$") -> None:
assert type(baseline) is type(candidate), f"type mismatch at {path}: {type(baseline)} != {type(candidate)}"
if isinstance(baseline, BaseModel) and isinstance(candidate, BaseModel):
assert_value_parity(_public_model_values(baseline), _public_model_values(candidate), path=path)
return
if isinstance(baseline, Mapping) and isinstance(candidate, Mapping):
baseline_mapping: Final = cast(Mapping[object, object], baseline)
candidate_mapping: Final = cast(Mapping[object, object], candidate)
assert frozenset((type(key), key) for key in baseline_mapping) == frozenset(
(type(key), key) for key in candidate_mapping
), f"mapping keys differ at {path}"
for key in baseline_mapping:
assert_value_parity(baseline_mapping[key], candidate_mapping[key], path=f"{path}.{key}")
return
if (
isinstance(baseline, Sequence)
and not isinstance(baseline, (str, bytes))
and isinstance(candidate, Sequence)
and not isinstance(candidate, (str, bytes))
):
baseline_sequence: Final = cast(Sequence[object], baseline)
candidate_sequence: Final = cast(Sequence[object], candidate)
assert len(baseline_sequence) == len(candidate_sequence), f"sequence lengths differ at {path}"
for index, (baseline_item, candidate_item) in enumerate(
zip(baseline_sequence, candidate_sequence, strict=True)
):
assert_value_parity(baseline_item, candidate_item, path=f"{path}[{index}]")
return
assert baseline == candidate, f"value mismatch at {path}: {baseline!r} != {candidate!r}"
def assert_parity(baseline: Execution, candidate: Execution, baseline_user_agent: str) -> None:
validate_harness(baseline, candidate, baseline_user_agent)
assert_request_parity(baseline.requests, candidate.requests)
assert_value_parity(baseline.report, candidate.report)

View file

@ -0,0 +1,64 @@
from __future__ import annotations
from collections.abc import Mapping
from typing import ClassVar, Final, Generic, Literal, TypeVar, cast
from pydantic import BaseModel, ConfigDict, Field, JsonValue, model_validator
from .recorded_http import RecordedResponse
JsonObject = dict[str, JsonValue]
class FixtureModel(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid", populate_by_name=True, serialize_by_alias=True)
class SdkInputBase(FixtureModel):
fixture_only_fields: ClassVar[tuple[str, ...]] = ()
def as_sdk_kwargs(self) -> dict[str, object]:
return cast(
dict[str, object],
self.model_dump(
mode="python",
exclude_unset=True,
exclude=set(self.fixture_only_fields),
),
)
def canonical_input(self) -> dict[str, object]:
dumped: Final = cast(dict[str, object], self.model_dump(mode="json", exclude_unset=True))
fixture_fields: Final = {field: getattr(self, field) for field in self.fixture_only_fields}
return {**fixture_fields, **dumped}
class JsonSchemaDefinition(FixtureModel):
name: str
description: str | None = None
schema_definition: JsonObject = Field(alias="schema")
strict: bool = False
class JsonSchemaResponseFormat(FixtureModel):
type: Literal["json_schema"]
json_schema: JsonSchemaDefinition
InputT = TypeVar("InputT", bound=SdkInputBase)
class ParityCase(FixtureModel, Generic[InputT]):
litellm_input: InputT
provider_responses: tuple[RecordedResponse, ...]
@model_validator(mode="before")
@classmethod
def load_legacy_single_response(cls, value: object) -> object:
if not isinstance(value, Mapping):
return value
migrated: Final = dict(cast(Mapping[str, object], value))
provider_response: Final = migrated.pop("provider_response", None)
if "provider_responses" not in migrated and provider_response is not None:
migrated["provider_responses"] = (provider_response,)
return migrated

View file

@ -0,0 +1 @@
from __future__ import annotations

View file

@ -0,0 +1,147 @@
from __future__ import annotations
from collections.abc import Mapping
from datetime import datetime
from itertools import accumulate
from typing import Final, Literal
from pydantic import AwareDatetime, BaseModel, ConfigDict, Field, TypeAdapter
from vcr.serialize import serialize
from vcr.serializers import yamlserializer
from .recording import RecordedInteraction
from ..recorded_http import (
HttpHeader,
RecordedHttpResponse,
RecordedHttpStreamResponse,
RecordedResponse,
RecordedStreamChunk,
)
_OBJECT: Final = TypeAdapter(dict[str, object])
class _CassetteModel(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid", populate_by_name=True)
class _Body(_CassetteModel):
string: str | bytes
def as_bytes(self) -> bytes:
return self.string.encode("utf-8") if isinstance(self.string, str) else self.string
class _Status(_CassetteModel):
code: int
message: str
class _Request(_CassetteModel):
method: str
uri: str
body: str | bytes | None
headers: dict[str, tuple[str, ...]]
class _Response(_CassetteModel):
status: _Status
headers: dict[str, tuple[str, ...]]
body: _Body
chunk_lengths: tuple[int, ...] | None = Field(default=None, alias="x-litellm-chunk-lengths")
def recorded_response(self) -> RecordedResponse:
headers: Final = tuple(
HttpHeader(name=name, value=value) for name, values in self.headers.items() for value in values
)
body: Final = self.body.as_bytes()
if self.chunk_lengths is None:
return RecordedHttpResponse.from_bytes(self.status.code, headers, body)
if any(length < 0 for length in self.chunk_lengths) or sum(self.chunk_lengths) != len(body):
raise ValueError("cassette stream chunk lengths do not match the response body")
offsets: Final = tuple(accumulate(self.chunk_lengths, initial=0))
return RecordedHttpStreamResponse(
kind="http_stream",
status_code=self.status.code,
headers=headers,
chunks=tuple(RecordedStreamChunk.from_bytes(body[start:end]) for start, end in zip(offsets, offsets[1:])),
)
class _Interaction(_CassetteModel):
request: _Request
response: _Response
class _ParityMetadata(_CassetteModel):
schema_version: Literal[1]
request_source: Literal["recorded", "python_replay"]
case: dict[str, object]
class ParityCassette(_CassetteModel):
version: Literal[1]
recorded_at: AwareDatetime
ttl_seconds: Literal[0]
interactions: tuple[_Interaction, ...]
parity: _ParityMetadata = Field(alias="x-litellm")
def case_data(self) -> dict[str, object]:
return {
**self.parity.case,
"provider_responses": tuple(item.response.recorded_response() for item in self.interactions),
}
def _response_dict(response: RecordedResponse) -> dict[str, object]:
headers: Final = {
name: [header.value for header in response.headers if header.name == name]
for name in dict.fromkeys(header.name for header in response.headers)
}
chunks: Final = (
tuple(chunk.data_bytes() for chunk in response.chunks)
if isinstance(response, RecordedHttpStreamResponse)
else None
)
body: Final = response.body_bytes() if isinstance(response, RecordedHttpResponse) else b"".join(chunks or ())
return {
"status": {"code": response.status_code, "message": ""},
"headers": headers,
"body": {"string": body},
**({"x-litellm-chunk-lengths": list(map(len, chunks))} if chunks is not None else {}),
}
def serialize_cassette(
case: Mapping[str, object],
interactions: tuple[RecordedInteraction, ...],
recorded_at: datetime,
request_source: Literal["recorded", "python_replay"],
) -> str:
normalized: Final = _OBJECT.validate_python(
yamlserializer.deserialize(
serialize(
{
"requests": [item.request for item in interactions],
"responses": [_response_dict(item.response) for item in interactions],
},
yamlserializer,
)
)
)
payload: Final = {
**normalized,
"recorded_at": recorded_at.isoformat(),
"ttl_seconds": 0,
"x-litellm": {
"schema_version": 1,
"request_source": request_source,
"case": {key: value for key, value in case.items() if key != "provider_responses"},
},
}
ParityCassette.model_validate(payload).case_data()
return str(yamlserializer.serialize(payload))
def deserialize_cassette(contents: str) -> ParityCassette:
return ParityCassette.model_validate(yamlserializer.deserialize(contents))

View file

@ -0,0 +1,34 @@
from __future__ import annotations
import argparse
from collections.abc import Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Final, cast
@dataclass(frozen=True, slots=True)
class RecordingArgs:
concurrency: int
examples: int
fixture_dir: Path | None
def _positive_int(value: str) -> int:
parsed: Final = int(value)
if parsed < 1:
raise argparse.ArgumentTypeError("must be at least 1")
return parsed
def parse_recording_args(argv: Sequence[str] | None = None) -> RecordingArgs:
parser: Final = argparse.ArgumentParser()
parser.add_argument("--concurrency", type=_positive_int, default=2)
parser.add_argument("--examples", type=_positive_int, default=4)
parser.add_argument("--fixture-dir", type=Path)
namespace: Final = parser.parse_args(argv)
return RecordingArgs(
concurrency=cast(int, namespace.concurrency),
examples=cast(int, namespace.examples),
fixture_dir=cast(Path | None, namespace.fixture_dir),
)

View file

@ -0,0 +1,22 @@
from __future__ import annotations
import queue
from typing import Final, TypeVar
from hypothesis import given, settings
from hypothesis.strategies import SearchStrategy
InputT = TypeVar("InputT")
def generate_case_inputs(strategy: SearchStrategy[InputT], examples: int) -> tuple[InputT, ...]:
generated: Final[queue.SimpleQueue[InputT | None]] = queue.SimpleQueue()
@settings(max_examples=examples, deadline=None, derandomize=True)
@given(case_input=strategy)
def generate_case(case_input: InputT) -> None:
generated.put(case_input)
generate_case()
generated.put(None)
return tuple(iter(generated.get, None))

View file

@ -0,0 +1,225 @@
from __future__ import annotations
import base64
from functools import cache
from io import BytesIO
from typing import Final
from urllib.parse import quote
from PIL import Image, ImageDraw
from reportlab.graphics.barcode import code128 # pyright: ignore[reportMissingTypeStubs] # ReportLab has no stubs
from reportlab.lib import colors # pyright: ignore[reportMissingTypeStubs] # ReportLab has no stubs
from reportlab.lib.pagesizes import letter # pyright: ignore[reportMissingTypeStubs] # ReportLab has no stubs
from reportlab.lib.utils import ImageReader # pyright: ignore[reportMissingTypeStubs] # ReportLab has no stubs
from reportlab.pdfgen import canvas # pyright: ignore[reportMissingTypeStubs] # ReportLab has no stubs
def dummy_image_url(text: str, font_size: int, width: int = 800, height: int = 300) -> str:
return f"https://dummyjson.com/image/{width}x{height}/ffffff/000000?text={quote(text)}&fontSize={font_size}"
_GLYPHS: Final = {
"D": ("11110", "10001", "10001", "10001", "10001", "10001", "11110"),
"O": ("01110", "10001", "10001", "10001", "10001", "10001", "01110"),
"C": ("01111", "10000", "10000", "10000", "10000", "10000", "01111"),
"1": ("00100", "01100", "00100", "00100", "00100", "00100", "01110"),
"2": ("01110", "10001", "00001", "00010", "00100", "01000", "11111"),
"3": ("11110", "00001", "00001", "01110", "00001", "00001", "11110"),
}
@cache
def structured_image_bytes() -> bytes:
image: Final = Image.new("RGB", (320, 80), "white")
draw: Final = ImageDraw.Draw(image)
scale: Final = 8
cursor_x = 24
for character in "DOC 123":
if character == " ":
cursor_x += scale * 3
continue
for glyph_y, row in enumerate(_GLYPHS[character]):
for glyph_x, filled in enumerate(row):
if filled == "1":
x = cursor_x + glyph_x * scale
y = 12 + glyph_y * scale
draw.rectangle((x, y, x + scale - 1, y + scale - 1), fill="black")
cursor_x += scale * 6
output: Final = BytesIO()
image.save(output, format="PNG")
return output.getvalue()
@cache
def structured_image_data_uri() -> str:
encoded: Final = base64.b64encode(structured_image_bytes()).decode("ascii")
return f"data:image/png;base64,{encoded}"
def _draw_header(pdf: canvas.Canvas, title: str, page_number: int) -> None:
pdf.setFillColor(colors.black)
pdf.setFont("Helvetica", 11)
pdf.drawString(45, 770, "Quarterly Operations Report")
pdf.setFont("Helvetica-Bold", 16)
pdf.drawString(45, 745, title)
pdf.setFont("Helvetica", 9)
pdf.drawString(45, 30, f"Confidential | Page {page_number} of 5")
def _draw_body(pdf: canvas.Canvas, page_number: int) -> None:
pdf.setFont("Helvetica", 10)
for line_number in range(1, 9):
pdf.drawString(
45,
500 - (line_number * 28),
f"Section {page_number}.{line_number}: Invoice totals, regional revenue, and reconciliation notes.",
)
def _diagram_image(width: int, height: int, accent: tuple[int, int, int]) -> Image.Image:
image: Final = Image.new("RGB", (width, height), (242, 246, 252))
draw: Final = ImageDraw.Draw(image)
for coordinate in range(0, max(width, height), 40):
draw.line((coordinate, 0, coordinate, height), fill=(32, 32, 32), width=3)
draw.line((0, coordinate, width, coordinate), fill=(32, 32, 32), width=3)
draw.line((0, 0, width, height), fill=accent, width=8)
draw.line((width, 0, 0, height), fill=accent, width=8)
draw.rectangle((width // 4, height // 4, width * 3 // 4, height * 3 // 4), outline=accent, width=6)
return image
def _draw_embedded_images(pdf: canvas.Canvas) -> None:
images: Final = (
(_diagram_image(320, 320, (51, 115, 217)), 455, 655, 70, 70),
(_diagram_image(360, 320, (38, 151, 92)), 455, 565, 70, 62),
(_diagram_image(120, 120, (219, 68, 55)), 455, 500, 45, 45),
)
for image, x, y, width, height in images:
pdf.drawImage( # pyright: ignore[reportUnknownMemberType] # ReportLab has no stubs
ImageReader(image), x, y, width=width, height=height, mask="auto"
)
def _draw_table_page(pdf: canvas.Canvas) -> None:
columns: Final = (45, 245, 405, 565)
tables: Final = (
(
(730, 695, 660, 625),
(
("Item", "Quantity", "Amount", 707),
("Document analysis", "2", "120.00", 672),
("Document verification", "1", "80.00", 637),
),
),
(
(600, 565, 530, 495),
(
("Item continued", "Quantity", "Amount", 577),
("Fixture validation", "3", "45.00", 542),
("Provider review", "1", "25.00", 507),
),
),
)
for rows, values in tables:
for x in columns:
pdf.line(x, rows[-1], x, rows[0])
for y in rows:
pdf.line(45, y, 565, y)
for item, quantity, amount, y in values:
pdf.drawString(55, y, item)
pdf.drawString(255, y, quantity)
pdf.drawString(415, y, amount)
def _draw_chart_page(pdf: canvas.Canvas) -> None:
bars: Final = ((70, 70), (170, 115), (270, 90), (370, 130))
pdf.setFillColor(colors.HexColor("#3373D9"))
for x, height in bars:
pdf.rect(x, 610, 65, height, fill=1, stroke=0)
pdf.setFillColor(colors.black)
for quarter, x in zip(("Q1", "Q2", "Q3", "Q4"), (90, 190, 290, 390), strict=True):
pdf.drawString(x, 590, quarter)
pdf.drawString(45, 550, "Formula: gross margin = (revenue - cost) / revenue")
_draw_embedded_images(pdf)
def _draw_metadata_page(pdf: canvas.Canvas) -> None:
pdf.setFont("Helvetica", 12)
pdf.drawString(45, 700, "Invoice Number: INV-2048")
pdf.drawString(45, 675, "Purchase Order: PO-4096")
pdf.setFillColor(colors.HexColor("#F2E65A"))
pdf.rect(40, 555, 500, 24, fill=1, stroke=0)
pdf.setFillColor(colors.black)
pdf.drawString(45, 560, "Highlighted total requiring review")
pdf.drawString(45, 530, "Reviewer comment: verify the highlighted total before approval")
pdf.setFillColor(colors.red)
pdf.drawString(45, 495, "Revised total: 245.00")
pdf.line(45, 501, 150, 501)
pdf.setFillColor(colors.black)
pdf.linkURL( # pyright: ignore[reportUnknownMemberType] # ReportLab has no stubs
"https://example.com/invoices/INV-2048", (45, 575, 300, 590), relative=0
)
pdf.highlightAnnotation( # pyright: ignore[reportUnknownMemberType] # ReportLab has no stubs
"Total highlighted for review",
Rect=(40, 555, 540, 579),
QuadPoints=(40, 579, 540, 579, 40, 555, 540, 555),
)
pdf.textAnnotation( # pyright: ignore[reportUnknownMemberType] # ReportLab has no stubs
"Verify the highlighted total", Rect=(520, 525, 540, 545)
)
pdf.drawString(45, 575, "https://example.com/invoices/INV-2048")
barcode: Final = code128.Code128("5901234123457", barHeight=70, barWidth=1.2)
barcode.drawOn(pdf, 90, 130)
def _draw_signature_page(pdf: canvas.Canvas) -> None:
pdf.saveState()
pdf.setFillColor(colors.lightgrey)
pdf.setFont("Helvetica-Bold", 54)
pdf.translate(110, 390)
pdf.rotate(25)
pdf.drawString(0, 0, "DRAFT")
pdf.restoreState()
pdf.setFillColor(colors.black)
pdf.setFont("Helvetica", 12)
pdf.drawString(45, 635, "Approved by: Jordan Lee")
pdf.line(45, 610, 310, 610)
pdf.bezier(55, 595, 75, 625, 112, 602, 155, 600)
pdf.drawString(45, 580, "Signature")
def _draw_appendix_page(pdf: canvas.Canvas) -> None:
pdf.setFont("Helvetica-Bold", 14)
pdf.drawString(45, 700, "1. Scope")
pdf.drawString(45, 650, "2. Findings")
pdf.drawString(45, 600, "3. Recommendations")
@cache
def structured_pdf_bytes() -> bytes:
output: Final = BytesIO()
pdf: Final = canvas.Canvas(output, pagesize=letter, pageCompression=0, invariant=1)
pdf.setTitle("Quarterly Operations Report")
pdf.setAuthor("LiteLLM parity fixture generator")
pdf.setSubject("Semantic document coverage for tables, figures, annotations, and metadata")
pdf.setKeywords("document, invoice, table, figure, annotation")
pages: Final = (
("Invoice Summary and Line Items", _draw_table_page),
("Revenue Chart and Formula Review", _draw_chart_page),
("Key Values, Link, Highlight, and Comment", _draw_metadata_page),
("Approval Signature and Watermark", _draw_signature_page),
("Appendix with Section Boundaries", _draw_appendix_page),
)
for page_number, (title, draw_page) in enumerate(pages, start=1):
_draw_header(pdf, title, page_number)
draw_page(pdf)
_draw_body(pdf, page_number)
pdf.showPage()
pdf.save()
return output.getvalue()
@cache
def structured_pdf_data_uri() -> str:
encoded: Final = base64.b64encode(structured_pdf_bytes()).decode("ascii")
return f"data:application/pdf;base64,{encoded}"

View file

@ -0,0 +1,201 @@
from __future__ import annotations
import logging
from concurrent.futures import Future, ThreadPoolExecutor, as_completed
from dataclasses import dataclass, field
from pathlib import Path
from types import MappingProxyType
from typing import Final, Generic, Literal, Protocol, TypeVar
from hypothesis.strategies import SearchStrategy
from pydantic import BaseModel
from .inputs import generate_case_inputs
from .recording import UpstreamEndpoint, record_upstream_interactions
from .store import (
FixtureInput,
canonical_json,
fixture_cache_key,
fixture_id,
fixture_path,
load_fixture,
save_fixture,
)
LOGGER: Final = logging.getLogger(__name__)
InputT = TypeVar("InputT", bound=FixtureInput)
InputT_contra = TypeVar("InputT_contra", bound=FixtureInput, contravariant=True)
CaseT = TypeVar("CaseT", bound=BaseModel)
class RecordingInvocation(Protocol[InputT_contra]):
def execute(self, provider_url: str, case_input: InputT_contra) -> None: ...
@dataclass(frozen=True, slots=True)
class RecordingTarget(Generic[InputT]):
name: str
upstream: UpstreamEndpoint
strategy: SearchStrategy[InputT]
invocation: RecordingInvocation[InputT] = field(repr=False)
required_inputs: tuple[InputT, ...] = ()
@dataclass(frozen=True, slots=True)
class RecordingJob(Generic[InputT]):
target_name: str
directory: Path
upstream: UpstreamEndpoint
case_input: InputT
invocation: RecordingInvocation[InputT] = field(repr=False)
@property
def case_id(self) -> str:
return fixture_id(self.case_input, self.target_name)
@dataclass(frozen=True, slots=True)
class RecordedFixture:
target_name: str
case_id: str
path: Path
kind: Literal["recorded"] = field(default="recorded", init=False)
@dataclass(frozen=True, slots=True)
class CachedFixture:
target_name: str
case_id: str
path: Path
kind: Literal["cached"] = field(default="cached", init=False)
@dataclass(frozen=True, slots=True)
class FailedFixture:
target_name: str
case_id: str
error: Exception = field(repr=False)
kind: Literal["failed"] = field(default="failed", init=False)
RecordingOutcome = RecordedFixture | CachedFixture | FailedFixture
@dataclass(frozen=True, slots=True)
class RecordingSummary:
recorded: tuple[RecordedFixture, ...]
cached: tuple[CachedFixture, ...]
failed: tuple[FailedFixture, ...]
@property
def exit_code(self) -> int:
return 1 if self.failed else 0
def _unique_inputs(target: RecordingTarget[InputT], examples: int) -> tuple[InputT, ...]:
generated_inputs: Final = generate_case_inputs(target.strategy, examples)
case_inputs: Final = (*target.required_inputs, *generated_inputs)
return tuple({canonical_json(fixture_cache_key(case_input)): case_input for case_input in case_inputs}.values())
def build_recording_jobs(
targets: tuple[RecordingTarget[InputT], ...],
root: Path,
examples: int,
) -> tuple[RecordingJob[InputT], ...]:
if examples < 1:
raise ValueError("examples must be at least 1")
return tuple(
RecordingJob(
target_name=target.name,
directory=root / target.name,
upstream=target.upstream,
case_input=case_input,
invocation=target.invocation,
)
for target in targets
for case_input in _unique_inputs(target, examples)
)
def _record_job(job: RecordingJob[InputT], case_type: type[CaseT]) -> RecordedFixture | CachedFixture:
cached: Final = load_fixture(job.directory, job.case_input, case_type)
if cached is not None:
path: Final = fixture_path(job.directory, job.case_input)
return CachedFixture(
target_name=job.target_name,
case_id=job.case_id,
path=path if path.is_file() else path.with_suffix(".json"),
)
interactions: Final = record_upstream_interactions(
job.upstream,
job.case_input,
job.invocation.execute,
)
status: Final = interactions[-1].response.status_code
if status in {408, 429} or status >= 500:
raise RuntimeError(f"Upstream returned transient HTTP {status}; rerun recording to retry")
case: Final = case_type.model_validate(
{
"litellm_input": job.case_input,
"provider_responses": tuple(item.response for item in interactions),
}
)
saved_path: Final = save_fixture(job.directory, job.case_input, case, interactions)
return RecordedFixture(target_name=job.target_name, case_id=job.case_id, path=saved_path)
def _completed_outcome(
completed: int,
total: int,
job: RecordingJob[InputT],
future: Future[RecordedFixture | CachedFixture],
) -> RecordingOutcome:
try:
outcome: Final = future.result()
except Exception as error:
failed: Final = FailedFixture(target_name=job.target_name, case_id=job.case_id, error=error)
LOGGER.error(
"[%d/%d] failed %s %s: %s",
completed,
total,
failed.target_name,
failed.case_id,
type(error).__name__,
)
return failed
LOGGER.info("[%d/%d] %s %s %s", completed, total, outcome.kind, outcome.target_name, outcome.case_id)
return outcome
def record_fixtures(
targets: tuple[RecordingTarget[InputT], ...],
root: Path,
examples: int,
concurrency: int,
case_type: type[CaseT],
) -> RecordingSummary:
if concurrency < 1:
raise ValueError("concurrency must be at least 1")
jobs: Final = build_recording_jobs(targets, root, examples)
total: Final = len(jobs)
LOGGER.info("Recording %d fixtures across %d targets with concurrency %d", total, len(targets), concurrency)
with ThreadPoolExecutor(max_workers=concurrency) as executor:
future_jobs: Final = MappingProxyType({executor.submit(_record_job, job, case_type): job for job in jobs})
outcomes: Final = tuple(
_completed_outcome(completed, total, future_jobs[future], future)
for completed, future in enumerate(as_completed(future_jobs), start=1)
)
summary: Final = RecordingSummary(
recorded=tuple(outcome for outcome in outcomes if isinstance(outcome, RecordedFixture)),
cached=tuple(outcome for outcome in outcomes if isinstance(outcome, CachedFixture)),
failed=tuple(outcome for outcome in outcomes if isinstance(outcome, FailedFixture)),
)
LOGGER.info(
"Finished %d fixtures: %d recorded, %d cached, %d failed",
total,
len(summary.recorded),
len(summary.cached),
len(summary.failed),
)
return summary

View file

@ -0,0 +1,66 @@
from __future__ import annotations
import os
from collections.abc import Callable
from pathlib import Path
from typing import Final, TypeVar
import pytest
from pydantic import BaseModel, ValidationError
from .store import recorded_fixtures
CaseT = TypeVar("CaseT", bound=BaseModel)
def parametrize_recorded_fixtures(
metafunc: pytest.Metafunc,
*,
fixture_name: str,
case_type: type[CaseT],
env_var: str,
default_directory: Path,
regeneration_command: str,
id_builder: Callable[[CaseT], str],
marks_builder: Callable[[CaseT], tuple[pytest.MarkDecorator, ...]] | None = None,
) -> None:
if fixture_name not in metafunc.fixturenames:
return
configured: Final = os.environ.get(env_var)
if configured == "":
raise pytest.UsageError(f"{env_var} is set but empty")
directory: Final = Path(configured).expanduser() if configured is not None else default_directory
try:
fixtures: Final = recorded_fixtures(directory, case_type)
except (ValidationError, ValueError) as error:
raise pytest.UsageError(
f"Invalid parity fixture bundle at {directory}. "
"Each fixture must use the current versioned envelope. "
f"Record fresh fixtures in an empty directory with: `{regeneration_command}`. "
f"Validation details: {error}"
) from error
if fixtures:
metafunc.parametrize(
fixture_name,
tuple(
pytest.param(
fixture,
id=id_builder(fixture),
marks=marks_builder(fixture) if marks_builder is not None else (),
)
for fixture in fixtures
),
)
return
if configured is not None:
raise pytest.UsageError(f"no recorded fixtures in {directory}")
metafunc.parametrize(
fixture_name,
(
pytest.param(
None,
marks=pytest.mark.skip(reason=f"no recorded fixtures in {directory}"),
id="no-recorded-fixtures",
),
),
)

View file

@ -0,0 +1,267 @@
from __future__ import annotations
import queue
import threading
from collections.abc import Callable, Generator, Iterable
from contextlib import contextmanager
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Final, TypeVar, cast
from urllib.parse import urlsplit, urlunsplit
import httpx
from vcr.filters import remove_query_parameters
from vcr.request import Request
from ..http import (
dropped_request_headers,
dropped_response_headers,
is_streaming_response,
)
from ..recorded_http import (
HttpHeader,
RecordedHttpResponse,
RecordedHttpStreamResponse,
RecordedResponse,
RecordedStreamChunk,
)
_PARITY_PROVIDER_HOST: Final = "parity-provider.invalid"
_SECRET_HEADERS: Final = frozenset(
{
"authorization",
"proxy-authorization",
"cookie",
"x-api-key",
"api-key",
"anthropic-api-key",
"openai-api-key",
"azure-api-key",
"x-goog-api-key",
"ocp-apim-subscription-key",
"x-amz-security-token",
}
)
InputT = TypeVar("InputT")
@dataclass(frozen=True, slots=True)
class UpstreamEndpoint:
base_url: str
@dataclass(frozen=True, slots=True)
class RecordedInteraction:
request: Request
response: RecordedResponse
def _end_to_end_headers(headers: httpx.Headers) -> tuple[HttpHeader, ...]:
decoded: Final = tuple((name.decode("ascii"), value.decode("latin-1")) for name, value in headers.raw)
excluded: Final = dropped_response_headers(decoded)
return tuple(
HttpHeader(name=name, value=_normalized_response_header(name, value))
for name, value in decoded
if name.lower() not in excluded
)
def _normalized_response_header(name: str, value: str) -> str:
if name.lower() not in {"location", "operation-location"}:
return value
parsed: Final = urlsplit(value)
if not parsed.netloc:
return value
return urlunsplit(("http", _PARITY_PROVIDER_HOST, parsed.path, parsed.query, parsed.fragment))
def local_response_header(name: str, value: str, provider_url: str) -> str:
if name.lower() not in {"location", "operation-location"}:
return value
parsed: Final = urlsplit(value)
if parsed.hostname != _PARITY_PROVIDER_HOST:
return value
return f"{provider_url}{parsed.path}{'?' + parsed.query if parsed.query else ''}"
class _RecordingProvider(ThreadingHTTPServer):
daemon_threads = True
def __init__(self, spec: UpstreamEndpoint) -> None:
super().__init__(("127.0.0.1", 0), _RecordingHandler)
self.spec: Final = spec
self.interactions: queue.Queue[RecordedInteraction] = queue.Queue()
@property
def url(self) -> str:
return f"http://127.0.0.1:{self.server_address[1]}"
def take_interactions(self) -> tuple[RecordedInteraction, ...]:
try:
first: Final = self.interactions.get(timeout=5)
except queue.Empty as error:
raise RuntimeError("successful SDK call did not produce a recorded response") from error
remaining: Final = tuple(self.interactions.get_nowait() for _ in range(self.interactions.qsize()))
return (first, *remaining)
class _RecordingHandler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def do_POST(self) -> None:
self._forward()
def do_GET(self) -> None:
self._forward()
def do_PUT(self) -> None:
self._forward()
def do_PATCH(self) -> None:
self._forward()
def do_DELETE(self) -> None:
self._forward()
def _forward(self) -> None:
provider: Final = self.server
assert isinstance(provider, _RecordingProvider)
length: Final = int(self.headers.get("content-length") or "0")
request_body: Final = self.rfile.read(length) if length else b""
raw_headers: Final = tuple(self.headers.raw_items())
excluded: Final = dropped_request_headers(raw_headers)
forwarded_headers: Final = tuple((name, value) for name, value in raw_headers if name.lower() not in excluded)
upstream_url: Final = f"{provider.spec.base_url.rstrip('/')}{self.path}"
try:
with httpx.stream(
self.command,
upstream_url,
headers=forwarded_headers,
content=request_body,
timeout=120,
) as upstream:
headers: Final = _end_to_end_headers(upstream.headers)
recorded_response: Final = self._record_upstream_response(upstream, headers)
except httpx.HTTPError as error:
self._send_response(502, (), str(error).encode("utf-8"))
return
recorded_request: Final = remove_query_parameters(
Request(
self.command,
f"http://{_PARITY_PROVIDER_HOST}{self.path}",
request_body,
{name: value for name, value in forwarded_headers if name.lower() not in _SECRET_HEADERS},
),
("api_key", "api-key", "key", "access_token", "subscription-key"),
)
provider.interactions.put(RecordedInteraction(recorded_request, recorded_response))
if isinstance(recorded_response, RecordedHttpResponse):
self._send_response(
recorded_response.status_code, recorded_response.headers, recorded_response.body_bytes()
)
def _record_upstream_response(
self,
upstream: httpx.Response,
headers: tuple[HttpHeader, ...],
) -> RecordedResponse:
content_type: Final = cast(str, upstream.headers.get("content-type", ""))
if is_streaming_response(content_type):
return self._record_stream(upstream, headers)
response_body: Final = b"".join(upstream.iter_bytes())
return RecordedHttpResponse.from_bytes(
status_code=upstream.status_code,
headers=headers,
body=response_body,
)
def _record_stream(
self,
upstream: httpx.Response,
headers: tuple[HttpHeader, ...],
) -> RecordedHttpStreamResponse:
self.send_response_only(upstream.status_code)
provider: Final = self.server
assert isinstance(provider, _RecordingProvider)
for header in headers:
self.send_header(header.name, local_response_header(header.name, header.value, provider.url))
self.send_header("transfer-encoding", "chunked")
self.end_headers()
chunks: Final = tuple(self._relay_chunks(upstream.iter_bytes()))
self.wfile.write(b"0\r\n\r\n")
self.wfile.flush()
return RecordedHttpStreamResponse(
kind="http_stream",
status_code=upstream.status_code,
headers=headers,
chunks=chunks,
)
def _relay_chunks(self, chunks: Iterable[bytes]) -> Generator[RecordedStreamChunk, None, None]:
for chunk in chunks:
self.wfile.write(f"{len(chunk):X}\r\n".encode("ascii"))
self.wfile.write(chunk)
self.wfile.write(b"\r\n")
self.wfile.flush()
yield RecordedStreamChunk.from_bytes(chunk)
def _send_response(self, status_code: int, headers: tuple[HttpHeader, ...], body: bytes) -> None:
self.send_response_only(status_code)
provider: Final = self.server
assert isinstance(provider, _RecordingProvider)
for header in headers:
self.send_header(header.name, local_response_header(header.name, header.value, provider.url))
self.send_header("content-length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def log_message(self, format: str, *args: object) -> None:
return
@contextmanager
def _recording_provider(spec: UpstreamEndpoint) -> Generator[_RecordingProvider]:
server: Final = _RecordingProvider(spec)
thread: Final = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield server
finally:
server.shutdown()
server.server_close()
thread.join(timeout=5)
def _invoke_and_take_interactions(
recorder: _RecordingProvider,
case_input: InputT,
sdk_call: Callable[[str, InputT], object],
) -> tuple[RecordedInteraction, ...]:
try:
sdk_call(recorder.url, case_input)
except Exception as invocation_error:
try:
return recorder.take_interactions()
except RuntimeError:
raise invocation_error
return recorder.take_interactions()
def record_upstream_interactions(
spec: UpstreamEndpoint,
case_input: InputT,
sdk_call: Callable[[str, InputT], object],
) -> tuple[RecordedInteraction, ...]:
with _recording_provider(spec) as recorder:
return _invoke_and_take_interactions(recorder, case_input, sdk_call)
def record_upstream_responses(
spec: UpstreamEndpoint,
case_input: InputT,
sdk_call: Callable[[str, InputT], object],
) -> tuple[RecordedResponse, ...]:
return tuple(item.response for item in record_upstream_interactions(spec, case_input, sdk_call))

View file

@ -0,0 +1,126 @@
from __future__ import annotations
import hashlib
import json
import tempfile
from collections.abc import Mapping
from datetime import datetime, timezone
from pathlib import Path
from typing import Final, Literal, Protocol, TypeVar, cast
from pydantic import AwareDatetime, BaseModel, ConfigDict, TypeAdapter, ValidationError
from .cassette import deserialize_cassette, serialize_cassette
from .recording import RecordedInteraction
FIXTURE_SCHEMA_VERSION: Final = 1
JSON_OBJECT: Final = TypeAdapter(dict[str, object])
class FixtureInput(Protocol):
def canonical_input(self) -> dict[str, object]: ...
CaseT = TypeVar("CaseT", bound=BaseModel)
class FixtureEnvelope(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
schema_version: int
recorded_at: AwareDatetime
case: dict[str, object]
def canonical_json(value: Mapping[str, object]) -> str:
return json.dumps(value, sort_keys=True, separators=(",", ":"), ensure_ascii=True)
def fixture_cache_key(case_input: FixtureInput) -> dict[str, object]:
return case_input.canonical_input()
def fixture_path(directory: Path, case_input: FixtureInput) -> Path:
input_json: Final = canonical_json(fixture_cache_key(case_input))
digest: Final = hashlib.sha256(input_json.encode("utf-8")).hexdigest()
return directory / f"{digest}.yaml"
def load_fixture(directory: Path, case_input: FixtureInput, case_type: type[CaseT]) -> CaseT | None:
path: Final = fixture_path(directory, case_input)
if path.is_file():
return read_fixture(path, case_type)
legacy_path: Final = path.with_suffix(".json")
if not legacy_path.is_file():
return None
return read_fixture(legacy_path, case_type)
def save_fixture(
directory: Path,
case_input: FixtureInput,
case: BaseModel,
interactions: tuple[RecordedInteraction, ...],
*,
recorded_at: datetime | None = None,
request_source: Literal["recorded", "python_replay"] = "recorded",
) -> Path:
directory.mkdir(parents=True, exist_ok=True)
path: Final = fixture_path(directory, case_input)
serialized: Final = serialize_cassette(
cast(dict[str, object], case.model_dump(mode="json", exclude_unset=True)),
interactions,
recorded_at or datetime.now(timezone.utc),
request_source,
)
with tempfile.NamedTemporaryFile(mode="w", encoding="utf-8", dir=directory, delete=False) as temporary:
temporary_path: Final = Path(temporary.name)
try:
temporary.write(serialized)
temporary.close()
temporary_path.replace(path)
finally:
temporary_path.unlink(missing_ok=True)
return path
def read_fixture(path: Path, case_type: type[CaseT]) -> CaseT:
contents: Final = path.read_text(encoding="utf-8")
if path.suffix == ".json":
return _load_fixture(JSON_OBJECT.validate_json(contents), path, case_type)
try:
cassette: Final = deserialize_cassette(contents)
return case_type.model_validate(cassette.case_data())
except ValueError as error:
raise ValueError(f"invalid parity cassette {path}") from error
def _load_fixture(raw_fixture: dict[str, object], path: Path, case_type: type[CaseT]) -> CaseT:
schema_version: Final = raw_fixture.get("schema_version")
if schema_version != FIXTURE_SCHEMA_VERSION:
raise ValueError(
f"fixture {path} has schema_version {schema_version!r}, expected {FIXTURE_SCHEMA_VERSION}; "
"delete it and regenerate the fixture bundle"
)
try:
envelope: Final = FixtureEnvelope.model_validate(raw_fixture)
return case_type.model_validate(envelope.case)
except ValidationError as error:
raise ValueError(f"invalid parity fixture {path} ({len(error.errors())} validation errors)") from error
def recorded_fixtures(directory: Path, case_type: type[CaseT]) -> tuple[CaseT, ...]:
if not directory.is_dir():
return ()
paths: Final = tuple(sorted((*directory.rglob("*.yaml"), *directory.rglob("*.json"))))
return tuple(read_fixture(path, case_type) for path in paths)
def fixture_directory(configured: Path | None, env_value: str | None, default: Path) -> Path:
return (configured or Path(env_value or default)).expanduser()
def fixture_id(case_input: FixtureInput, prefix: str) -> str:
input_json: Final = canonical_json(case_input.canonical_input())
digest: Final = hashlib.sha256(input_json.encode("utf-8")).hexdigest()[:8]
return f"{prefix}-{digest}"

View file

@ -0,0 +1,95 @@
from __future__ import annotations
from datetime import datetime, timezone
from pathlib import Path
from typing import Final
import httpx
import pytest
from vcr import VCR
from vcr.request import Request
from ..fixture_models import ParityCase, SdkInputBase
from .cassette import deserialize_cassette
from .recording import RecordedInteraction
from .store import load_fixture, save_fixture
from ..recorded_http import (
HttpHeader,
RecordedHttpResponse,
RecordedHttpStreamResponse,
RecordedResponse,
RecordedStreamChunk,
)
from ..replay import replay_server
_URI: Final = "http://parity-provider.invalid/operation?api-version=1"
class _Input(SdkInputBase):
model: str = "fixture-model"
@pytest.mark.parametrize("body", (b'{"text":"caf\xc3\xa9"}', b"\x00\xff\x80", b""))
def test_cassette_replays_repeated_requests_with_vcr_and_preserves_bytes(tmp_path: Path, body: bytes) -> None:
sdk_input: Final = _Input()
responses: Final = tuple(
RecordedHttpResponse.from_bytes(
status,
(HttpHeader(name="content-type", value="application/octet-stream"),),
body,
)
for status in (200, 429)
)
case: Final = ParityCase[_Input](litellm_input=sdk_input, provider_responses=responses)
interactions: Final = tuple(
RecordedInteraction(Request("POST", _URI, b"\xffrequest", {}), response) for response in responses
)
timestamp: Final = datetime(2020, 1, 1, tzinfo=timezone.utc)
path: Final = save_fixture(tmp_path, sdk_input, case, interactions, recorded_at=timestamp)
assert load_fixture(tmp_path, sdk_input, ParityCase[_Input]) == case
assert deserialize_cassette(path.read_text()).recorded_at == timestamp
with VCR().use_cassette(str(path), record_mode="none", match_on=("method", "uri", "body")) as cassette:
for status in (200, 429):
replayed: Final = httpx.post(_URI, content=b"\xffrequest")
assert replayed.status_code == status
assert replayed.content == body
assert cassette.all_played
def test_stream_cassette_preserves_chunk_boundaries_through_local_replay(tmp_path: Path) -> None:
sdk_input: Final = _Input()
chunks: Final = (b"data: caf\xc3", b"\xa9\n\n", b"data: [DONE]\n\n")
response: Final = RecordedHttpStreamResponse(
kind="http_stream",
status_code=200,
headers=(HttpHeader(name="content-type", value="text/event-stream"),),
chunks=tuple(RecordedStreamChunk.from_bytes(chunk) for chunk in chunks),
)
case: Final = ParityCase[_Input](litellm_input=sdk_input, provider_responses=(response,))
path: Final = save_fixture(
tmp_path, sdk_input, case, (RecordedInteraction(Request("POST", _URI, b"{}", {}), response),)
)
loaded: Final = load_fixture(tmp_path, sdk_input, ParityCase[_Input])
assert loaded == case
with replay_server() as server:
server.enqueue_response(loaded.provider_responses[0])
with httpx.stream("POST", f"{server.url}/operation", content=b"{}") as replayed:
assert tuple(replayed.iter_raw()) == chunks
server.take_requests(1)
path.write_text(path.read_text().replace("- 10\n", "- 999\n"))
with pytest.raises(ValueError, match="invalid parity cassette"):
load_fixture(tmp_path, sdk_input, ParityCase[_Input])
def test_cassette_preserves_duplicate_response_headers(tmp_path: Path) -> None:
sdk_input: Final = _Input()
response: Final[RecordedResponse] = RecordedHttpResponse.from_bytes(
200,
(HttpHeader(name="x-test", value="first"), HttpHeader(name="x-test", value="second")),
b"{}",
)
case: Final = ParityCase[_Input](litellm_input=sdk_input, provider_responses=(response,))
save_fixture(tmp_path, sdk_input, case, (RecordedInteraction(Request("POST", _URI, b"", {}), response),))
assert load_fixture(tmp_path, sdk_input, ParityCase[_Input]) == case

View file

@ -0,0 +1,20 @@
from __future__ import annotations
from typing import Final
from hypothesis import strategies as st
from pydantic import BaseModel, ConfigDict
from .inputs import generate_case_inputs
class _Input(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
identifier: str
def test_generate_case_inputs_is_deterministic() -> None:
strategy: Final = st.builds(_Input, identifier=st.integers().map(str))
assert generate_case_inputs(strategy, examples=4) == generate_case_inputs(strategy, examples=4)

View file

@ -0,0 +1,27 @@
from __future__ import annotations
import base64
from io import BytesIO
from typing import Final, cast
from PIL import Image
from .media import dummy_image_url, structured_image_bytes, structured_image_data_uri
def test_dummy_image_url_encodes_text_and_dimensions() -> None:
assert dummy_image_url("invoice 123", 24, width=320, height=80) == (
"https://dummyjson.com/image/320x80/ffffff/000000?text=invoice%20123&fontSize=24"
)
def test_structured_image_is_local_content_bearing_png() -> None:
png: Final = structured_image_bytes()
encoded: Final = structured_image_data_uri().partition(",")[2]
image: Final = Image.open(BytesIO(png))
colors: Final = cast(list[tuple[int, tuple[int, int, int]]], image.getcolors(maxcolors=2))
assert png.startswith(b"\x89PNG\r\n\x1a\n")
assert base64.b64decode(encoded, validate=True) == png
assert image.size == (320, 80)
assert {color for _, color in colors} == {(0, 0, 0), (255, 255, 255)}

View file

@ -0,0 +1,217 @@
from __future__ import annotations
import logging
import threading
from collections.abc import Generator
from contextlib import contextmanager
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Final, Literal
import httpx
import pytest
from hypothesis import strategies as st
from pydantic import BaseModel, ConfigDict
from .pipeline import (
RecordingInvocation,
RecordingTarget,
build_recording_jobs,
record_fixtures,
)
from .recording import UpstreamEndpoint
from .store import fixture_path
from ..recorded_http import RecordedResponse
class _FixtureInput(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
identifier: str
def canonical_input(self) -> dict[str, object]:
return {"identifier": self.identifier}
class _ParityCase(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
litellm_input: _FixtureInput
provider_responses: tuple[RecordedResponse, ...]
class _Upstream(ThreadingHTTPServer):
daemon_threads = True
def __init__(self, status: int = 200) -> None:
super().__init__(("127.0.0.1", 0), _UpstreamHandler)
self.response_status: Final = status
@property
def url(self) -> str:
return f"http://127.0.0.1:{self.server_address[1]}"
class _UpstreamHandler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def do_POST(self) -> None:
length: Final = int(self.headers.get("content-length") or "0")
self.rfile.read(length)
body: Final = b"{}"
server: Final = self.server
assert isinstance(server, _Upstream)
self.send_response(server.response_status)
self.send_header("content-type", "application/json")
self.send_header("content-length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def log_message(self, format: str, *args: object) -> None:
return
@contextmanager
def _upstream(status: int = 200) -> Generator[_Upstream]:
server: Final = _Upstream(status)
thread: Final = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield server
finally:
server.shutdown()
server.server_close()
thread.join(timeout=5)
@dataclass(frozen=True, slots=True)
class _OrderedInvocation:
order: Literal["slow", "fast"]
slow_started: threading.Event
fast_finished: threading.Event
def execute(self, provider_url: str, case_input: _FixtureInput) -> None:
if self.order == "slow":
self.slow_started.set()
if not self.fast_finished.wait(timeout=2):
raise TimeoutError("fast recording did not finish")
else:
if not self.slow_started.wait(timeout=2):
raise TimeoutError("slow recording did not start")
response: Final = httpx.post(f"{provider_url}/record", json={"id": case_input.identifier}, timeout=5)
response.raise_for_status()
if self.order == "fast":
self.fast_finished.set()
@dataclass(frozen=True, slots=True)
class _Invocation:
def execute(self, provider_url: str, case_input: _FixtureInput) -> None:
response: Final = httpx.post(f"{provider_url}/record", json={"id": case_input.identifier}, timeout=5)
response.raise_for_status()
def _target(
name: str,
upstream_url: str,
case_input: _FixtureInput,
invocation: RecordingInvocation[_FixtureInput],
) -> RecordingTarget[_FixtureInput]:
return RecordingTarget(
name=name,
upstream=UpstreamEndpoint(base_url=upstream_url),
strategy=st.just(case_input),
invocation=invocation,
required_inputs=(case_input,),
)
def test_build_jobs_keeps_required_inputs_before_generated_inputs_and_deduplicates(tmp_path: Path) -> None:
required: Final = _FixtureInput(identifier="required")
generated: Final = _FixtureInput(identifier="generated")
target: Final = RecordingTarget(
name="ordered",
upstream=UpstreamEndpoint(base_url="https://provider.invalid"),
strategy=st.just(generated),
invocation=_Invocation(),
required_inputs=(required, required),
)
jobs: Final = build_recording_jobs((target,), tmp_path, examples=1)
assert tuple(job.case_input.identifier for job in jobs) == ("required", "generated")
def test_progress_follows_completion_order(tmp_path: Path, caplog: pytest.LogCaptureFixture) -> None:
slow_started: Final = threading.Event()
fast_finished: Final = threading.Event()
with _upstream() as upstream:
targets: Final = (
_target(
"slow",
upstream.url,
_FixtureInput(identifier="slow"),
_OrderedInvocation("slow", slow_started, fast_finished),
),
_target(
"fast",
upstream.url,
_FixtureInput(identifier="fast"),
_OrderedInvocation("fast", slow_started, fast_finished),
),
)
with caplog.at_level(logging.INFO, logger="tests.rust-python-harness.shared.parity.fixtures.pipeline"):
summary: Final = record_fixtures(targets, tmp_path, 1, 2, _ParityCase)
progress: Final = tuple(record.message for record in caplog.records if record.message.startswith("["))
assert len(summary.recorded) == 2
assert summary.exit_code == 0
assert "recorded fast" in progress[0]
assert "recorded slow" in progress[1]
assert caplog.records[0].message == "Recording 2 fixtures across 2 targets with concurrency 2"
assert caplog.records[-1].message == "Finished 2 fixtures: 2 recorded, 0 cached, 0 failed"
def test_failure_does_not_stop_independent_recordings(tmp_path: Path) -> None:
stale_input: Final = _FixtureInput(identifier="stale")
stale_directory: Final = tmp_path / "stale"
stale_directory.mkdir()
fixture_path(stale_directory, stale_input).with_suffix(".json").write_text(
'{"schema_version": 0}\n', encoding="utf-8"
)
with _upstream() as upstream:
targets: Final = (
_target("stale", upstream.url, stale_input, _Invocation()),
_target("valid", upstream.url, _FixtureInput(identifier="valid"), _Invocation()),
)
summary: Final = record_fixtures(targets, tmp_path, 1, 2, _ParityCase)
assert len(summary.recorded) == 1
assert summary.recorded[0].target_name == "valid"
assert len(summary.failed) == 1
assert summary.failed[0].target_name == "stale"
assert summary.exit_code == 1
@pytest.mark.parametrize("status", (408, 429, 500, 503))
def test_transient_response_is_not_cached_and_can_be_retried(tmp_path: Path, status: int) -> None:
case_input: Final = _FixtureInput(identifier="retry")
with _upstream(status) as upstream:
target: Final = _target("retry", upstream.url, case_input, _Invocation())
failed: Final = record_fixtures((target,), tmp_path, 1, 1, _ParityCase)
assert failed.exit_code == 1
assert not fixture_path(tmp_path / "retry", case_input).exists()
with _upstream() as healthy_upstream:
healthy_target: Final = _target("retry", healthy_upstream.url, case_input, _Invocation())
retried: Final = record_fixtures((healthy_target,), tmp_path, 1, 1, _ParityCase)
assert retried.exit_code == 0
assert len(retried.recorded) == 1
def test_provider_rejected_response_can_be_recorded(tmp_path: Path) -> None:
with _upstream(400) as upstream:
target: Final = _target("rejected", upstream.url, _FixtureInput(identifier="invalid"), _Invocation())
summary: Final = record_fixtures((target,), tmp_path, 1, 1, _ParityCase)
assert summary.exit_code == 0
assert len(summary.recorded) == 1

View file

@ -0,0 +1,574 @@
from __future__ import annotations
import asyncio
import queue
import threading
from collections.abc import AsyncIterator, Callable, Generator, Iterator
from contextlib import contextmanager
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Final, Literal
import httpx
import pytest
from hypothesis import strategies as st
from openai._streaming import SSEDecoder
from pydantic import BaseModel, ConfigDict
from ..compare import assert_request_parity
from .pipeline import RecordingTarget, record_fixtures
from .recording import (
UpstreamEndpoint,
record_upstream_interactions,
record_upstream_responses,
)
from .store import (
FIXTURE_SCHEMA_VERSION,
fixture_path,
load_fixture,
recorded_fixtures,
)
from ..inprocess import InProcessExecution, run_in_process, run_in_process_async
from ..recorded_http import (
HttpHeader,
RecordedHttpStreamResponse,
RecordedResponse,
RecordedStreamChunk,
)
from ..replay import ReplayServer, replay_server
from ..stream import (
StreamCompleted,
StreamFailed,
StreamOutcome,
assert_stream_parity,
consume_async_stream,
consume_sync_stream,
)
_SSE_CHUNKS: Final = (
b'data: {"choices":[{"delta":{"content":"hello"}}]}\n\n',
b'data: {"choices":[{"delta":{"content":" world"}}]}\n\n',
b"data: [DONE]\n\n",
)
class _StreamEvent(BaseModel):
kind: Literal["delta", "done", "error"]
value: str
class _StreamApplicationError(Exception):
status_code: Final = 400
code: Final = "invalid_input"
type: Final = "validation_error"
param: Final = "input"
model: Final = "fixture-model"
llm_provider: Final = "fixture-provider"
def _stream_event(data: str) -> _StreamEvent:
event: Final = _StreamEvent.model_validate_json(data)
if event.kind == "error":
raise _StreamApplicationError(event.value)
return event
def _event_chunks(failed: bool) -> tuple[bytes, ...]:
terminal: Final = (
b'event: error\r\ndata: {"kind":"error","value":"invalid input"}\r\n\r\n'
if failed
else b'event: done\r\ndata: {"kind":"done","value":""}\r\n\r\n'
)
return (
b'event: delta\r\ndata: {"kind":"delta",\r\ndata: "value":"caf\xc3',
b'\xa9"}\r\n',
b'\r\nevent: delta\r\ndata: {"kind":"delta","value":"second"}\r\n\r\n' + terminal,
)
def _sync_events(api_base: str, case_input: _FixtureInput) -> Iterator[_StreamEvent]:
with httpx.stream("POST", f"{api_base}/stream", json={"id": case_input.identifier}, timeout=5) as response:
response.raise_for_status()
for event in SSEDecoder().iter_bytes(response.iter_bytes()):
yield _stream_event(event.data)
async def _async_events(api_base: str, case_input: _FixtureInput) -> AsyncIterator[_StreamEvent]:
async with httpx.AsyncClient(timeout=5) as client:
async with client.stream("POST", f"{api_base}/stream", json={"id": case_input.identifier}) as response:
response.raise_for_status()
async for event in SSEDecoder().aiter_bytes(response.aiter_bytes()):
yield _stream_event(event.data)
async def _consume_async_events(api_base: str, case_input: _FixtureInput) -> StreamOutcome:
async def create() -> AsyncIterator[_StreamEvent]:
return _async_events(api_base, case_input)
return await consume_async_stream(create)
async def _replay_events(
mode: Literal["sync", "async"],
provider: ReplayServer,
response: RecordedHttpStreamResponse,
case_input: _FixtureInput,
) -> InProcessExecution[StreamOutcome]:
if mode == "sync":
return run_in_process(
provider, (response,), lambda url: consume_sync_stream(lambda: _sync_events(url, case_input))
)
return await run_in_process_async(provider, (response,), lambda url: _consume_async_events(url, case_input))
class _FixtureInput(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
identifier: str
def canonical_input(self) -> dict[str, object]:
return {"identifier": self.identifier}
class _ParityCase(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
litellm_input: _FixtureInput
provider_responses: tuple[RecordedResponse, ...]
@dataclass(frozen=True, slots=True)
class _Invocation:
sdk_call: Callable[[str, _FixtureInput], object]
def execute(self, provider_url: str, case_input: _FixtureInput) -> None:
self.sdk_call(provider_url, case_input)
class _ControlledUpstream(ThreadingHTTPServer):
daemon_threads = True
def __init__(self, stream_chunks: tuple[bytes, ...]) -> None:
super().__init__(("127.0.0.1", 0), _ControlledUpstreamHandler)
self.stream_chunks: Final = stream_chunks
self.lock: Final = threading.Lock()
self.two_requests_started: Final = threading.Event()
self.active_requests: int = 0
self.max_active_requests: int = 0
self.request_count: int = 0
@property
def url(self) -> str:
return f"http://127.0.0.1:{self.server_address[1]}"
def start_request(self) -> None:
with self.lock:
self.active_requests += 1
self.request_count += 1
self.max_active_requests = max(self.max_active_requests, self.active_requests)
if self.active_requests == 2:
self.two_requests_started.set()
self.two_requests_started.wait(timeout=2)
def end_tracked_request(self) -> None:
with self.lock:
self.active_requests -= 1
class _ControlledUpstreamHandler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def do_POST(self) -> None:
upstream: Final = self.server
assert isinstance(upstream, _ControlledUpstream)
length: Final = int(self.headers.get("content-length") or "0")
self.rfile.read(length)
if self.path == "/credentials?api_key=query-secret&api-version=1":
authorized: Final = self.headers.get("authorization") == "Bearer header-secret"
self._send_json(200 if authorized else 401, b"{}")
return
if self.path == "/upload":
self._send_json(200, b'{"file_id":"fixture://document.pdf"}')
return
if self.path == "/parse":
self._send_json(200, b'{"result":{"chunks":[]}}')
return
if self.path == "/analyze":
self.send_response(202)
self.send_header("operation-location", f"{upstream.url}/results/1")
self.send_header("content-length", "0")
self.end_headers()
return
if self.path in {"/v1/chat/completions", "/stream"}:
with upstream.lock:
upstream.request_count += 1
self.send_response(200)
self.send_header("content-type", "text/event-stream")
self.send_header("transfer-encoding", "chunked")
self.end_headers()
for chunk in upstream.stream_chunks:
self.wfile.write(f"{len(chunk):X}\r\n".encode("ascii"))
self.wfile.write(chunk)
self.wfile.write(b"\r\n")
self.wfile.flush()
self.wfile.write(b"0\r\n\r\n")
self.wfile.flush()
return
if self.path == "/error":
self._send_json(429, b'{"error":{"message":"rate limited"}}')
return
upstream.start_request()
try:
body: Final = b"{}"
self.send_response(200)
self.send_header("content-type", "application/json")
self.send_header("set-cookie", "session=must-not-be-recorded")
self.send_header("content-length", str(len(body)))
self.end_headers()
self.wfile.write(body)
finally:
upstream.end_tracked_request()
def do_GET(self) -> None:
if self.path == "/results/1":
self._send_json(200, b'{"status":"succeeded","analyzeResult":{"pages":[]}}')
return
self.send_error(404)
def do_PUT(self) -> None:
self.do_POST()
def do_PATCH(self) -> None:
self.do_POST()
def do_DELETE(self) -> None:
self.do_POST()
def _send_json(self, status: int, body: bytes) -> None:
self.send_response(status)
self.send_header("content-type", "application/json")
self.send_header("content-length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def log_message(self, format: str, *args: object) -> None:
return
@contextmanager
def _controlled_upstream(stream_chunks: tuple[bytes, ...] = _SSE_CHUNKS) -> Generator[_ControlledUpstream]:
server: Final = _ControlledUpstream(stream_chunks)
thread: Final = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield server
finally:
server.shutdown()
server.server_close()
thread.join(timeout=5)
def _case(identifier: str) -> _FixtureInput:
return _FixtureInput(identifier=identifier)
def _sdk_call(api_base: str, case_input: _FixtureInput) -> object:
return httpx.post(f"{api_base}/v1/operation", content=b"{}", timeout=5)
def _stream_sdk_call(api_base: str, case_input: _FixtureInput) -> object:
return httpx.post(f"{api_base}/v1/chat/completions", content=b"{}", timeout=5)
def _error_sdk_call(api_base: str, case_input: _FixtureInput) -> object:
response: Final = httpx.post(f"{api_base}/error", content=b"{}", timeout=5)
response.raise_for_status()
return response
def _method_sdk_call(method: str) -> Callable[[str, _FixtureInput], object]:
def call(api_base: str, case_input: _FixtureInput) -> object:
return httpx.request(method, f"{api_base}/method", json={"id": case_input.identifier}, timeout=5)
return call
def _multi_sdk_call(api_base: str, case_input: _FixtureInput) -> object:
upload: Final = httpx.post(f"{api_base}/upload", json={"document": case_input.identifier}, timeout=5)
upload.raise_for_status()
parsed: Final = httpx.post(f"{api_base}/parse", json={"input": upload.json()["file_id"]}, timeout=5)
parsed.raise_for_status()
return parsed
def _polling_sdk_call(api_base: str, case_input: _FixtureInput) -> object:
started: Final = httpx.post(f"{api_base}/analyze", json={"document": case_input.identifier}, timeout=5)
operation_location: Final = started.headers["operation-location"]
completed: Final = httpx.get(operation_location, timeout=5)
completed.raise_for_status()
return completed
def test_recording_deduplicates_per_target_and_caps_global_concurrency(tmp_path: Path) -> None:
shared_input: Final = _case("shared")
with _controlled_upstream() as upstream:
spec: Final = UpstreamEndpoint(base_url=upstream.url)
targets: Final = (
RecordingTarget(
name="first",
upstream=spec,
strategy=st.just(shared_input),
invocation=_Invocation(_sdk_call),
required_inputs=(shared_input, shared_input),
),
RecordingTarget(
name="second",
upstream=spec,
strategy=st.just(shared_input),
invocation=_Invocation(_sdk_call),
required_inputs=(shared_input,),
),
)
summary: Final = record_fixtures(targets, tmp_path, examples=1, concurrency=2, case_type=_ParityCase)
assert len(summary.recorded) == 2
assert {result.target_name for result in summary.recorded} == {"first", "second"}
assert summary.cached == ()
assert summary.failed == ()
assert upstream.request_count == 2
assert upstream.max_active_requests == 2
assert len(recorded_fixtures(tmp_path, _ParityCase)) == 2
for path in tmp_path.rglob("*.yaml"):
contents = path.read_text(encoding="utf-8")
assert f"schema_version: {FIXTURE_SCHEMA_VERSION}" in contents
assert "recorded_at:" in contents
def test_pipeline_rejects_stale_fixture_before_provider_call(tmp_path: Path) -> None:
case_input: Final = _case("stale")
directory: Final = tmp_path / "stale-target"
directory.mkdir()
path: Final = fixture_path(directory, case_input).with_suffix(".json")
path.write_text('{"schema_version": 0}\n', encoding="utf-8")
target: Final = RecordingTarget(
name="stale-target",
upstream=UpstreamEndpoint(base_url="http://127.0.0.1:1"),
strategy=st.just(case_input),
invocation=_Invocation(_sdk_call),
)
summary: Final = record_fixtures(
(target,),
tmp_path,
examples=1,
concurrency=1,
case_type=_ParityCase,
)
assert summary.recorded == ()
assert summary.cached == ()
assert len(summary.failed) == 1
assert str(summary.failed[0].error) == (
f"fixture {path} has schema_version 0, expected {FIXTURE_SCHEMA_VERSION}; "
"delete it and regenerate the fixture bundle"
)
def test_cached_fixture_is_reported_without_provider_call(tmp_path: Path) -> None:
case_input: Final = _case("cached")
with _controlled_upstream() as upstream:
target: Final = RecordingTarget(
name="cached-target",
upstream=UpstreamEndpoint(base_url=upstream.url),
strategy=st.just(case_input),
invocation=_Invocation(_sdk_call),
)
first: Final = record_fixtures((target,), tmp_path, 1, 1, _ParityCase)
second: Final = record_fixtures((target,), tmp_path, 1, 1, _ParityCase)
assert len(first.recorded) == 1
assert len(second.cached) == 1
assert upstream.request_count == 1
def test_streaming_response_records_and_replays_chunks() -> None:
with _controlled_upstream() as upstream:
responses: Final = record_upstream_responses(
UpstreamEndpoint(base_url=upstream.url),
_case("stream"),
_stream_sdk_call,
)
response: Final = responses[0]
assert isinstance(response, RecordedHttpStreamResponse)
assert tuple(chunk.data_bytes() for chunk in response.chunks) == _SSE_CHUNKS
assert isinstance(response.model_dump(mode="json")["chunks"], list)
with replay_server() as provider:
provider.enqueue_response(response)
with httpx.stream("POST", f"{provider.url}/v1/chat/completions", json={}) as replayed:
replayed_chunks: Final = tuple(replayed.iter_raw())
provider.take_requests(1)
assert replayed_chunks == _SSE_CHUNKS
def test_non_successful_provider_response_is_recorded() -> None:
with _controlled_upstream() as upstream:
responses: Final = record_upstream_responses(
UpstreamEndpoint(base_url=upstream.url),
_case("provider-error"),
_error_sdk_call,
)
response: Final = responses[0]
assert response.status_code == 429
def test_sensitive_response_headers_are_not_recorded() -> None:
with _controlled_upstream() as upstream:
responses: Final = record_upstream_responses(
UpstreamEndpoint(base_url=upstream.url),
_case("headers"),
_sdk_call,
)
assert all(header.name.lower() != "set-cookie" for header in responses[0].headers)
def test_recorded_requests_strip_credentials_without_changing_the_live_request() -> None:
def sdk_call(api_base: str, case_input: _FixtureInput) -> object:
return httpx.post(
f"{api_base}/credentials?api_key=query-secret&api-version=1",
headers={
"Authorization": "Bearer header-secret",
"Ocp-Apim-Subscription-Key": "azure-secret",
"Cookie": "session=cookie-secret",
"X-Test": case_input.identifier,
},
content=b"\xffdocument",
)
with _controlled_upstream() as upstream:
interactions: Final = record_upstream_interactions(
UpstreamEndpoint(upstream.url), _case("credentials"), sdk_call
)
interaction: Final = interactions[0]
assert interaction.response.status_code == 200
assert interaction.request.uri == "http://parity-provider.invalid/credentials?api-version=1"
assert interaction.request.body == b"\xffdocument"
assert interaction.request.headers["x-test"] == "credentials"
assert all(
header not in interaction.request.headers for header in ("authorization", "ocp-apim-subscription-key", "cookie")
)
@pytest.mark.parametrize("method", ("PUT", "PATCH", "DELETE"))
def test_recording_and_replay_support_mutating_http_methods(method: str) -> None:
sdk_call: Final = _method_sdk_call(method)
with _controlled_upstream() as upstream:
responses: Final = record_upstream_responses(
UpstreamEndpoint(base_url=upstream.url),
_case(method),
sdk_call,
)
with replay_server() as provider:
provider.enqueue_response(responses[0])
sdk_call(provider.url, _case(method))
requests: Final = provider.take_requests(1)
assert requests[0].method == method
def test_stream_response_model_rejects_buffered_body() -> None:
with pytest.raises(ValueError, match="Extra inputs are not permitted"):
RecordedHttpStreamResponse.model_validate(
{
"kind": "http_stream",
"status_code": 200,
"headers": [HttpHeader(name="content-type", value="text/event-stream")],
"chunks": [RecordedStreamChunk.from_bytes(b"data: [DONE]\n\n")],
"body_b64": "",
}
)
@pytest.mark.parametrize("sdk_call", (_multi_sdk_call, _polling_sdk_call))
def test_multiple_provider_calls_record_and_replay_in_order(
sdk_call: Callable[[str, _FixtureInput], object],
) -> None:
with _controlled_upstream() as upstream:
responses: Final = record_upstream_responses(
UpstreamEndpoint(base_url=upstream.url),
_case(sdk_call.__name__),
sdk_call,
)
assert len(responses) == 2
with replay_server() as provider:
for response in responses:
provider.enqueue_response(response)
sdk_call(provider.url, _case(sdk_call.__name__))
requests: Final = provider.take_requests(2)
assert len(requests) == 2
@pytest.mark.asyncio
@pytest.mark.parametrize("mode", ("sync", "async"))
@pytest.mark.parametrize("failed", (False, True), ids=("completed", "application-error"))
async def test_typed_stream_recording_cassette_replay_parity(
tmp_path: Path, mode: Literal["sync", "async"], failed: bool
) -> None:
case_input: Final = _case("typed-stream")
outcomes: Final[queue.SimpleQueue[StreamOutcome]] = queue.SimpleQueue()
def record(api_base: str, sdk_input: _FixtureInput) -> None:
outcome: Final = (
consume_sync_stream(lambda: _sync_events(api_base, sdk_input))
if mode == "sync"
else asyncio.run(_consume_async_events(api_base, sdk_input))
)
outcomes.put(outcome)
with _controlled_upstream(_event_chunks(failed)) as upstream:
target: Final = RecordingTarget(
name="stream",
upstream=UpstreamEndpoint(upstream.url),
strategy=st.just(case_input),
invocation=_Invocation(record),
)
summary: Final = record_fixtures((target,), tmp_path, 1, 1, _ParityCase)
assert summary.failed == ()
assert len(summary.recorded) == 1
recorded: Final = outcomes.get_nowait()
loaded: Final = load_fixture(tmp_path / "stream", case_input, _ParityCase)
assert loaded is not None
response: Final = loaded.provider_responses[0]
assert isinstance(response, RecordedHttpStreamResponse)
assert response.status_code == 200
wire_bytes: Final = b"".join(chunk.data_bytes() for chunk in response.chunks)
assert wire_bytes == b"".join(_event_chunks(failed))
coalesced: Final = response.model_copy(update={"chunks": (RecordedStreamChunk.from_bytes(wire_bytes),)})
with replay_server() as provider:
first: Final = await _replay_events(mode, provider, response, case_input)
second: Final = await _replay_events(mode, provider, coalesced, case_input)
assert_request_parity(first.requests, second.requests)
assert len(first.requests) == 1
assert first.requests[0].body == {"id": case_input.identifier}
assert_stream_parity(recorded, first.response)
assert_stream_parity(first.response, second.response)
expected: Final = (_StreamEvent(kind="delta", value="café"), _StreamEvent(kind="delta", value="second"))
assert first.response.chunks == (expected if failed else (*expected, _StreamEvent(kind="done", value="")))
if failed:
assert isinstance(first.response.terminal, StreamFailed)
assert first.response.terminal.phase == "iteration"
assert first.response.terminal.exception_type is _StreamApplicationError
assert first.response.terminal.error.code == "invalid_input"
assert first.response.terminal.error.message == "invalid input"
else:
assert first.response.terminal == StreamCompleted()

View file

@ -0,0 +1,54 @@
from __future__ import annotations
from collections.abc import Iterable
from typing import Final
HOP_BY_HOP_HEADERS: Final[frozenset[str]] = frozenset(
{
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailer",
"trailers",
"transfer-encoding",
"upgrade",
}
)
REQUEST_DROPPED_HEADERS: Final[frozenset[str]] = HOP_BY_HOP_HEADERS | {
"host",
"content-length",
"accept-encoding",
}
RESPONSE_DROPPED_HEADERS: Final[frozenset[str]] = HOP_BY_HOP_HEADERS | {
"content-encoding",
"content-length",
"set-cookie",
}
def connection_header_names(headers: Iterable[tuple[str, str]]) -> frozenset[str]:
return frozenset(
token.strip().lower()
for name, value in headers
if name.lower() == "connection"
for token in value.split(",")
if token.strip()
)
def dropped_request_headers(headers: Iterable[tuple[str, str]]) -> frozenset[str]:
materialized: Final = tuple(headers)
return REQUEST_DROPPED_HEADERS | connection_header_names(materialized)
def dropped_response_headers(headers: Iterable[tuple[str, str]]) -> frozenset[str]:
materialized: Final = tuple(headers)
return RESPONSE_DROPPED_HEADERS | connection_header_names(materialized)
def is_streaming_response(content_type: str) -> bool:
return "text/event-stream" in content_type.lower()

View file

@ -0,0 +1,47 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import Final, Generic, TypeVar
from .models import CapturedRequest
from .recorded_http import RecordedResponse
from .replay import ReplayServer
ResponseT = TypeVar("ResponseT")
@dataclass(frozen=True, slots=True)
class InProcessExecution(Generic[ResponseT]):
requests: tuple[CapturedRequest, ...]
response: ResponseT
def run_in_process(
provider: ReplayServer,
recorded_responses: tuple[RecordedResponse, ...],
call: Callable[[str], ResponseT],
) -> InProcessExecution[ResponseT]:
for recorded_response in recorded_responses:
provider.enqueue_response(recorded_response)
try:
response: Final = call(provider.url)
return InProcessExecution(requests=provider.take_requests(len(recorded_responses)), response=response)
except Exception:
provider.reset()
raise
async def run_in_process_async(
provider: ReplayServer,
recorded_responses: tuple[RecordedResponse, ...],
call: Callable[[str], Awaitable[ResponseT]],
) -> InProcessExecution[ResponseT]:
for recorded_response in recorded_responses:
provider.enqueue_response(recorded_response)
try:
response: Final = await call(provider.url)
return InProcessExecution(requests=provider.take_requests(len(recorded_responses)), response=response)
except Exception:
provider.reset()
raise

View file

@ -0,0 +1,145 @@
from __future__ import annotations
import base64
from typing import Annotated, Final, Literal, cast
from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter
class CapturedRequest(BaseModel):
model_config = ConfigDict(frozen=True)
method: str
path: str
headers: tuple[tuple[str, str], ...]
body: JsonValue
user_agent: str | None
class SDKSuccess(BaseModel):
model_config = ConfigDict(frozen=True)
status: Literal["ok"] = "ok"
response: JsonValue
class SDKError(BaseModel):
model_config = ConfigDict(frozen=True)
status: Literal["error"] = "error"
exception_type: str
message: str
status_code: int | None
code: str | None
error_type: str | None
param: str | None
model: str | None
llm_provider: str | None
class SDKJsonChunk(BaseModel):
model_config = ConfigDict(frozen=True)
kind: Literal["json"] = "json"
value: JsonValue
class SDKBytesChunk(BaseModel):
model_config = ConfigDict(frozen=True)
kind: Literal["bytes"] = "bytes"
data_b64: str
def data_bytes(self) -> bytes:
return base64.b64decode(self.data_b64, validate=True)
SDKChunk = Annotated[SDKJsonChunk | SDKBytesChunk, Field(discriminator="kind")]
class SDKStreamCompleted(BaseModel):
model_config = ConfigDict(frozen=True)
kind: Literal["completed"] = "completed"
class SDKStreamFailed(BaseModel):
model_config = ConfigDict(frozen=True)
kind: Literal["failed"] = "failed"
error: SDKError
SDKStreamTerminal = Annotated[SDKStreamCompleted | SDKStreamFailed, Field(discriminator="kind")]
class SDKStreamReport(BaseModel):
model_config = ConfigDict(frozen=True)
status: Literal["stream"] = "stream"
chunks: tuple[SDKChunk, ...]
terminal: SDKStreamTerminal
SDKReport = Annotated[SDKSuccess | SDKError | SDKStreamReport, Field(discriminator="status")]
JSON_VALUE_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
def sdk_chunk(value: object) -> SDKChunk:
if isinstance(value, bytes):
return SDKBytesChunk(data_b64=base64.b64encode(value).decode("ascii"))
if isinstance(value, BaseModel):
return SDKJsonChunk(value=JSON_VALUE_ADAPTER.validate_python(value.model_dump(mode="json")))
return SDKJsonChunk(value=JSON_VALUE_ADAPTER.validate_python(value))
def _string_attribute(error: Exception, name: str) -> str | None:
value: Final = cast(object | None, getattr(error, name, None))
return None if value is None else str(value)
def sdk_error_report(error: Exception) -> SDKError:
message, _, _ = str(error).partition("\nTraceback (most recent call last):")
raw_status_code: Final = cast(object | None, getattr(error, "status_code", None))
status_code: Final = raw_status_code if isinstance(raw_status_code, int) else None
return SDKError(
exception_type=f"{type(error).__module__}.{type(error).__qualname__}",
message=message.rstrip(),
status_code=status_code,
code=_string_attribute(error, "code"),
error_type=_string_attribute(error, "type"),
param=_string_attribute(error, "param"),
model=_string_attribute(error, "model"),
llm_provider=_string_attribute(error, "llm_provider"),
)
class Execution(BaseModel):
model_config = ConfigDict(frozen=True)
requests: tuple[CapturedRequest, ...]
report: SDKReport
class SDKCommand(BaseModel):
model_config = ConfigDict(frozen=True)
case_file: str
route: str
class WorkerSuccess(BaseModel):
model_config = ConfigDict(frozen=True)
status: Literal["ok"] = "ok"
report: SDKReport
class WorkerFailure(BaseModel):
model_config = ConfigDict(frozen=True)
status: Literal["error"] = "error"
error: str
WorkerResult = Annotated[WorkerSuccess | WorkerFailure, Field(discriminator="status")]

View file

@ -0,0 +1,63 @@
from __future__ import annotations
import base64
from typing import Annotated, Literal
from pydantic import BaseModel, ConfigDict, Field
class _RecordedHttpModel(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
class HttpHeader(_RecordedHttpModel):
name: str
value: str
class RecordedHttpResponse(_RecordedHttpModel):
kind: Literal["http"]
status_code: int
headers: tuple[HttpHeader, ...]
body_b64: str
@classmethod
def from_bytes(
cls,
status_code: int,
headers: tuple[HttpHeader, ...],
body: bytes,
) -> RecordedHttpResponse:
return cls(
kind="http",
status_code=status_code,
headers=headers,
body_b64=base64.b64encode(body).decode("ascii"),
)
def body_bytes(self) -> bytes:
return base64.b64decode(self.body_b64, validate=True)
class RecordedStreamChunk(_RecordedHttpModel):
data_b64: str
@classmethod
def from_bytes(cls, data: bytes) -> RecordedStreamChunk:
return cls(data_b64=base64.b64encode(data).decode("ascii"))
def data_bytes(self) -> bytes:
return base64.b64decode(self.data_b64, validate=True)
class RecordedHttpStreamResponse(_RecordedHttpModel):
kind: Literal["http_stream"]
status_code: int
headers: tuple[HttpHeader, ...]
chunks: tuple[RecordedStreamChunk, ...]
RecordedResponse = Annotated[
RecordedHttpResponse | RecordedHttpStreamResponse,
Field(discriminator="kind"),
]

View file

@ -0,0 +1,147 @@
from __future__ import annotations
import base64
import queue
import threading
from collections.abc import Generator
from contextlib import contextmanager
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Final
from pydantic import JsonValue, TypeAdapter
from .fixtures.recording import local_response_header
from .models import CapturedRequest
from .recorded_http import RecordedHttpResponse, RecordedHttpStreamResponse, RecordedResponse
JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
EXCLUDED_REQUEST_HEADERS: Final = frozenset(
{
"host",
"content-length",
"connection",
"accept-encoding",
"user-agent",
"x-litellm-parity-route",
}
)
EXCLUDED_RESPONSE_HEADERS: Final = frozenset({"content-length", "transfer-encoding", "connection"})
class ReplayServer(ThreadingHTTPServer):
daemon_threads = True
def __init__(self) -> None:
super().__init__(("127.0.0.1", 0), _ReplayHandler)
self.responses: queue.Queue[RecordedResponse] = queue.Queue()
self.requests: queue.Queue[CapturedRequest] = queue.Queue()
@property
def url(self) -> str:
return f"http://127.0.0.1:{self.server_address[1]}"
def enqueue_response(self, response: RecordedResponse) -> None:
self.responses.put(response)
def take_requests(self, expected_count: int) -> tuple[CapturedRequest, ...]:
request_count: Final = self.requests.qsize()
if request_count != expected_count:
raise AssertionError(f"expected exactly {expected_count} provider requests, received {request_count}")
return tuple(self.requests.get_nowait() for _ in range(request_count))
def reset(self) -> None:
while not self.responses.empty():
self.responses.get_nowait()
while not self.requests.empty():
self.requests.get_nowait()
class _ReplayHandler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def do_POST(self) -> None:
self._replay()
def do_GET(self) -> None:
self._replay()
def do_PUT(self) -> None:
self._replay()
def do_PATCH(self) -> None:
self._replay()
def do_DELETE(self) -> None:
self._replay()
def _replay(self) -> None:
provider: Final = self.server
assert isinstance(provider, ReplayServer)
length: Final = int(self.headers.get("content-length") or "0")
raw_body: Final = self.rfile.read(length) if length else b""
content_type: Final = self.headers.get("content-type", "")
body: Final = (
JSON_VALUE.validate_json(raw_body)
if raw_body and content_type.lower().startswith("application/json")
else base64.b64encode(raw_body).decode("ascii")
if raw_body
else None
)
headers: Final = tuple(
sorted(
(name.lower(), value)
for name, value in self.headers.raw_items()
if name.lower() not in EXCLUDED_REQUEST_HEADERS
)
)
provider.requests.put(
CapturedRequest(
method=self.command,
path=self.path,
headers=headers,
body=body,
user_agent=self.headers.get("user-agent"),
)
)
try:
response: Final = provider.responses.get(timeout=5)
except queue.Empty:
self.send_error(500, "no replay response queued")
return
self.send_response_only(response.status_code)
for header in response.headers:
if header.name.lower() not in EXCLUDED_RESPONSE_HEADERS:
self.send_header(header.name, local_response_header(header.name, header.value, provider.url))
if isinstance(response, RecordedHttpResponse):
response_body: Final = response.body_bytes()
self.send_header("content-length", str(len(response_body)))
self.end_headers()
self.wfile.write(response_body)
return
assert isinstance(response, RecordedHttpStreamResponse)
self.send_header("transfer-encoding", "chunked")
self.end_headers()
for chunk in response.chunks:
data = chunk.data_bytes()
self.wfile.write(f"{len(data):X}\r\n".encode("ascii"))
self.wfile.write(data)
self.wfile.write(b"\r\n")
self.wfile.flush()
self.wfile.write(b"0\r\n\r\n")
self.wfile.flush()
def log_message(self, format: str, *args: object) -> None:
return
@contextmanager
def replay_server() -> Generator[ReplayServer]:
server: Final = ReplayServer()
thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.01}, daemon=True)
thread.start()
try:
yield server
finally:
server.shutdown()
server.server_close()
thread.join(timeout=5)

View file

@ -0,0 +1,198 @@
from __future__ import annotations
import asyncio
import os
import subprocess
import sys
from collections import deque
from collections.abc import Callable, Generator
from concurrent.futures import ThreadPoolExecutor, TimeoutError
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final, TextIO, cast
from pydantic import TypeAdapter, ValidationError
from .models import (
Execution,
SDKCommand,
WorkerFailure,
WorkerResult,
WorkerSuccess,
)
from .recorded_http import RecordedResponse
from .replay import ReplayServer, replay_server
WORKER_RESULT_PREFIX: Final = "LITELLM_PARITY_RESULT "
WORKER_RESULT_ADAPTER: Final[TypeAdapter[WorkerResult]] = TypeAdapter(WorkerResult)
@dataclass(frozen=True, slots=True)
class SubprocessRunner:
entrypoint: Path
baseline_user_agent: str
route_label: str
def command(self, provider_url: str) -> tuple[str, ...]:
return (
sys.executable,
"-m",
".".join(
self.entrypoint.resolve().relative_to(Path(__file__).resolve().parents[4]).with_suffix("").parts
),
"--parity-worker",
provider_url,
)
@dataclass(frozen=True, slots=True)
class ExecutionVariant:
name: str
environment: tuple[tuple[str, str], ...]
class SubprocessWorker:
def __init__(self, runner: SubprocessRunner, provider: ReplayServer, variant: ExecutionVariant) -> None:
project_root: Final = str(Path(__file__).resolve().parents[4])
existing_pythonpath: Final = os.environ.get("PYTHONPATH")
env: Final = {
**os.environ,
**dict(variant.environment),
"LITELLM_USER_AGENT": runner.baseline_user_agent,
"PYTHONPATH": os.pathsep.join(path for path in (project_root, existing_pythonpath) if path),
}
self.mode: Final = variant.name
self.route_label: Final = runner.route_label
self.provider: Final = provider
self.process: Final = subprocess.Popen(
runner.command(provider.url),
stdin=subprocess.PIPE,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
bufsize=1,
env=env,
)
self.output_reader: Final = ThreadPoolExecutor(max_workers=1)
self.recent_output: Final[deque[str]] = deque(maxlen=100)
def execute(
self,
case_file: Path,
route: str,
responses: tuple[RecordedResponse, ...],
) -> Execution:
stdin: Final = self.process.stdin
if stdin is None or self.process.poll() is not None:
raise AssertionError(f"{self.mode} {self.route_label} worker exited before processing {case_file}")
for response in responses:
self.provider.enqueue_response(response)
command: Final = SDKCommand(case_file=str(case_file), route=route)
try:
stdin.write(f"{command.model_dump_json()}\n")
stdin.flush()
result: Final = self.output_reader.submit(self._read_result).result(timeout=60)
except TimeoutError as error:
self.provider.reset()
self.close()
raise AssertionError(
f"{self.mode} {self.route_label} worker timed out after 60s while processing {case_file}"
) from error
except AssertionError:
self.provider.reset()
raise
except (BrokenPipeError, OSError) as error:
self.provider.reset()
raise AssertionError(self._failure_message(f"worker pipe failed while processing {case_file}")) from error
if isinstance(result, WorkerFailure):
self.provider.reset()
raise AssertionError(
f"{self.mode} {self.route_label} worker failed while processing {case_file}:\n{result.error}"
)
assert isinstance(result, WorkerSuccess)
try:
return Execution(requests=self.provider.take_requests(len(responses)), report=result.report)
except AssertionError:
self.provider.reset()
raise
def _read_result(self) -> WorkerResult:
process_stdout: Final = self.process.stdout
if process_stdout is None:
raise AssertionError(self._failure_message("worker stdout is unavailable"))
stdout: Final = cast(TextIO, process_stdout)
line: Final = stdout.readline()
if not line:
raise AssertionError(self._failure_message("worker exited without returning a result"))
stripped: Final = line.rstrip()
if not stripped.startswith(WORKER_RESULT_PREFIX):
self.recent_output.append(stripped)
return self._read_result()
payload: Final = stripped.removeprefix(WORKER_RESULT_PREFIX)
try:
return WORKER_RESULT_ADAPTER.validate_json(payload)
except ValidationError as error:
raise AssertionError(self._failure_message("worker returned an invalid result")) from error
def _failure_message(self, message: str) -> str:
output: Final = "\n".join(self.recent_output)
prefix: Final = f"{self.mode} {self.route_label}"
return f"{prefix} {message}" if not output else f"{prefix} {message}\noutput:\n{output}"
def close(self) -> None:
stdin: Final = self.process.stdin
if stdin is not None and not stdin.closed:
stdin.close()
try:
self.process.wait(timeout=10)
except subprocess.TimeoutExpired:
self.process.terminate()
self.process.wait(timeout=10)
self.output_reader.shutdown(wait=True, cancel_futures=True)
@contextmanager
def execution_worker(
runner: SubprocessRunner,
variant: ExecutionVariant,
) -> Generator[SubprocessWorker]:
with replay_server() as provider:
worker: Final = SubprocessWorker(runner, provider, variant)
try:
yield worker
finally:
worker.close()
def run_execution(
worker: SubprocessWorker,
case_file: Path,
route: str,
responses: tuple[RecordedResponse, ...],
) -> Execution:
return worker.execute(case_file, route, responses)
@contextmanager
def execution_worker_pair(
runner: SubprocessRunner,
baseline: ExecutionVariant,
candidate: ExecutionVariant,
) -> Generator[tuple[SubprocessWorker, SubprocessWorker]]:
with execution_worker(runner, baseline) as baseline_worker:
with execution_worker(runner, candidate) as candidate_worker:
yield baseline_worker, candidate_worker
def parity_worker_main(
execute_command: Callable[[str, str, asyncio.AbstractEventLoop], WorkerResult],
mock_url: str,
) -> None:
event_loop: Final = asyncio.new_event_loop()
try:
for line in sys.stdin:
sys.stdout.write(f"{WORKER_RESULT_PREFIX}{execute_command(line, mock_url, event_loop).model_dump_json()}\n")
sys.stdout.flush()
finally:
event_loop.close()

Some files were not shown because too many files have changed in this diff Show more